diff --git a/.github/ISSUE_TEMPLATE.md b/.github/ISSUE_TEMPLATE.md deleted file mode 100644 index ef3b9bc736..0000000000 --- a/.github/ISSUE_TEMPLATE.md +++ /dev/null @@ -1,13 +0,0 @@ -One line summary of the issue here. - -### Expected behavior - -As concisely as possible, describe the expected behavior. - -### Actual behavior - -As concisely as possible, describe the observed behavior. - -### Steps to reproduce the behavior - -Please list all relevant steps to reproduce the observed behavior. \ No newline at end of file diff --git a/.github/PULL_REQUEST_TEMPLATE.md b/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md similarity index 93% rename from .github/PULL_REQUEST_TEMPLATE.md rename to .github/PULL_REQUEST_TEMPLATE/pull_request_template.md index bfb079c94c..6454011e26 100644 --- a/.github/PULL_REQUEST_TEMPLATE.md +++ b/.github/PULL_REQUEST_TEMPLATE/pull_request_template.md @@ -13,3 +13,5 @@ Result What will change as a result of your pull request? Note that sometimes this section is unnecessary because it is self-explanatory based on the solution. + +Closes twitter/util#PR_NUMBER diff --git a/.github/workflows/continuous-integration.yml b/.github/workflows/continuous-integration.yml new file mode 100644 index 0000000000..8c34d0eafb --- /dev/null +++ b/.github/workflows/continuous-integration.yml @@ -0,0 +1,55 @@ +name: continuous integration + +env: + JAVA_OPTS: "-Dsbt.log.noformat=true" + TRAVIS: "true" # pretend we're TravisCI + SBT_VERSION: 1.7.1 + BUILD_VERSION: "v1" # bump this if builds are failing due to a bad cache + +defaults: + run: + shell: bash +on: + push: + branches: + - develop + - release + pull_request: + +jobs: + build: + strategy: + fail-fast: false + matrix: + scala: [2.12.12, 2.13.6] + java: ['1.8', '1.11'] + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v2.3.4 + - uses: actions/setup-java@v1 + with: + java-version: ${{ matrix.java }} + - name: echo java version + run: java -Xmx32m -version + - name: echo javac version + run: javac -J-Xmx32m -version + - name: cache build dependencies + uses: actions/cache@v2 + env: + cache-name: cache-build-deps + with: + path: | + ~/.ivy2/cache + ~/.ivy2/local/com.twitter + ~/.sbt + key: ${{ runner.os }}-build-${{ env.BUILD_VERSION }}-${{ env.cache-name }}-${{ env.SBT_VERSION }}-${{ matrix.scala }}-${{ matrix.java }} + - name: update cache + run: | + if [ -f ~/.ivy2/cache ]; then + find ~/.ivy2/cache -name "ivydata-*.properties" -delete + fi + if [ -f ~/.sbt ]; then + find ~/.sbt -name "*.lock" -delete + fi + - name: test + run: ${{ format('./sbt ++{0} clean test', matrix.scala) }} diff --git a/.gitignore b/.gitignore index a8dfda5e8c..59f053bcff 100644 --- a/.gitignore +++ b/.gitignore @@ -1,18 +1,16 @@ -target/ -dist/ -project/boot/ -project/plugins/project/ -project/plugins/src_managed/ -*.log -*.tmproj -lib_managed/ -*.swp +*.class *.iml -.idea/ -.idea/* +*.ipr +*.log .DS_Store -.ensime -.ivyjars -sbt +.cache +.history +.idea/ +.lib +dist/ +lib_managed/ out/ +project/boot/ +project/plugins/project/ sbt-launch.jar +target/ diff --git a/.travis.yml b/.travis.yml deleted file mode 100644 index 85263e6ebc..0000000000 --- a/.travis.yml +++ /dev/null @@ -1,43 +0,0 @@ -# container-based build -sudo: false - -language: scala - -env: - - JAVA_OPTS="-DSKIP_FLAKY=true -Dsbt.log.noformat=true" - -# These directories are cached to S3 at the end of the build -cache: - directories: - - $HOME/.ivy2 - - $HOME/.sbt/boot/scala-$TRAVIS_SCALA_VERSION - - $HOME/.dodo - -before_cache: - # Tricks to avoid unnecessary cache updates - - find $HOME/.ivy2 -name "ivydata-*.properties" -delete - - find $HOME/.sbt -name "*.lock" -delete - -scala: - - 2.11.11 - - 2.12.4 - -jdk: - - oraclejdk8 - -notifications: - slack: - secure: E0knpocRWQT1YI4TgAZGxS7Xoe15wWZ5fbJsi0nfMhxHSlngrnxAknN60qovPa2J8R/opryX/0jwaUF2/5RPYVxf6ARNI428WkrltxctQeCugN8j5XqpBqXLR/+GPetqFe9WwF6ycg3bmThsXnap6/xNB75zMWj+2E1CYTqwpA0= - -before_script: - # default $SBT_OPTS is irrelevant to sbt lancher - - unset SBT_OPTS - - travis_retry ./sbt --error ++$TRAVIS_SCALA_VERSION update - -script: - - travis_retry ./sbt ++$TRAVIS_SCALA_VERSION clean coverage test coverageReport - - if [[ "$TRAVIS_SCALA_VERSION" == 2.11* ]]; then travis_retry ./sbt ++$TRAVIS_SCALA_VERSION "project util-intellij" coverage test coverageReport; fi - - ./sbt ++$TRAVIS_SCALA_VERSION coverageAggregate - -after_success: - - bash <(curl -s https://codecov.io/bash) diff --git a/BUILD.bazel b/BUILD.bazel new file mode 100644 index 0000000000..dbd3d06196 --- /dev/null +++ b/BUILD.bazel @@ -0,0 +1 @@ +# This prevents SQ query from grabbing //:all since it traverses up once to find a BUILD (DPB-14048) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index af32da932f..266e2e3281 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -1,5 +1,5 @@ .. Author notes: this file is formatted with restructured text - (http://docutils.sourceforge.net/docs/user/rst/quickstart.html). + (https://docutils.sourceforge.net/docs/user/rst/quickstart.html). The changelog style is adapted from Apache Lucene. Note that ``PHAB_ID=#`` and ``RB_ID=#`` correspond to associated messages in commits. @@ -7,6 +7,1205 @@ Note that ``PHAB_ID=#`` and ``RB_ID=#`` correspond to associated messages in com Unreleased ---------- +New Features +~~~~~~~~~~~~ + +* util-jvm: Add gc pause stats for all collector pools, including G1. ``PHAB_ID=D1176049`` +* util-stats: Expose dimensional metrics APIs and allow metrics with an indeterminate + identity to be exported through the Prometheus exporter. ``PHAB_ID=D1218090`` + + +24.5.0 +------ + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-app: When the application exits due to an error on startup, the error and + and stack trace are printed to stderr in addition to the existing stdout. ``PHAB_ID=D1116753`` +* util-jvm: Memory manager & pool metrics are labelled dimensionally rather than hierarchically + when exported as dimensional metrics. ``PHAB_ID=D1218090`` +* util-app: Enables a jvm shutdown hook to call app.quit(). ``PHAB_ID=D1236932`` + +23.11.0 +------- + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.14.2. ``PHAB_ID=D1049772`` + +* util: Bump version of Jackson to 2.14.3. ``PHAB_ID=D1069160`` + +* util-securty: Add `deserializeAndFilterOutInvalidCertificates` Which wraps + the `deserializeX509` call in a Try. ``PHAB_ID=D1107551`` + +22.12.0 +------- + +New Features +~~~~~~~~~~~~ + +* util-core: Introduce Future.fromCompletableFuture ``PHAB_ID=D940686`` + +* util-core: Introduce `Time.fromInstant` ``PHAB_ID=D987464`` + +* util-core: Introduce `TimeFormatter`, a replacement for `TimeFormat`. `TimeFormatter` uses + the modern `java.time.format.DateTimeFormatter` API instead of the `java.text.SimpleDateFormat` + with the main advantage being that it doesn't require synchronization on either formatting or + parsing, and is thus more performant and causes no contention. ``PHAB_ID=D987464`` + +* util-stats: `MetricBuilder` now has API to configure a format in which histograms are exported. + Use `.withHistogramFormat` to access it. ``PHAB_ID=D908189`` + +Deprecations +~~~~~~~~~~~~ +* util-core: `TimeFormat`, as well as the `Time.format` methods are deprecated in favor of using + `TimeFormatter`. ``PHAB_ID=D987464`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.14.1. ``PHAB_ID=D1025778`` + +* util-core: `Time.at`, `Time.toString`, and `Time.fromRss` now use `TimeFormatter` internally and + thus have better performance and do not cause contention through synchronization. `toString` always + uses `TimeFormatter` while `at` and `fromRss` only use it when run with Java 9 or later due to a + compatibility bug in Java 1.8. ``PHAB_ID=D987464`` + +22.7.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-stats: SourceRole subtypes have been moved into the SourceRole object to make their + relationship to the parent type more clear. ``PHAB_ID=D909744`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-jackson: Update Jackson library to version 2.13.3 ``PHAB_ID=D906005`` + +* util-jackson: Deserialized case classes with validation on optional fields shouldn't throw an error. + ``PHAB_ID=D915377`` + +* util: Update snakeyaml to 1.28 ``PHAB_ID=D930268`` + +22.4.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-stats: The metric instantiation methods have been removed from `MetricBuilder`. Use the methods on + `StatsReceiver` to instantiate metrics instead. ``PHAB_ID=D859036`` + +22.3.0 +------ + +Deprecations +~~~~~~~~~~~~ + +* util-stats: Deprecated methods on `MetricBuilder` for directly instantiating metrics so that we can + eventually remove the `statsReceiever` field from `MetricBuilder` and let it just be the source of + truth for defining a metric. ``PHAB_ID=D862244`` + +22.2.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-core: Added `Memoize.classValue` as a Scala-friendly API for `java.lang.ClassValue`. ``PHAB_ID=D825673`` + +* util-jvm: Register JVM expression including memory pool usages (including code cache, compressed class space, + eden space, sheap, metaspace, survivor space, and old gen) and open file descriptors count in StatsReceiver. + ``PHAB_ID=D820472`` + +* util-slf4j-jul-bridge: Add `Slf4jBridge` trait which can be mixed into extensions of `c.t.app.App` in + order to attempt installation of the `SLF4JBridgeHandler` via the `Slf4jBridgeUtility` in the constructor + of the `c.t.app.App` instance. ``PHAB_ID=D827493`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.13.2. ``PHAB_ID=D848592`` + +* util-slf4j-api: Update the `Logger` API to include "call-by-name" method + variations akin to the `Logging` trait. When creating a `Logger` from a + Scala singleton object class, the resultant logger name will no longer include + the `$` suffix. Remove the deprecated `Loggers` object which is no longer + needed for Java compatibility as users can now directly use the `Logger` + apply functions with no additional ceremony. ``PHAB_ID=D820252`` + +* util: Bump version of Caffeine to 2.9.3. ``PHAB_ID=D815761`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: The `c.t.util.Responder` trait has been removed. ``PHAB_ID=D824830`` + +22.1.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-jackson: The error message when failing to deserialize a character now correctly prints the non-character string. ``PHAB_ID=D808049`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.13.1. ``PHAB_ID=D808049`` + +* util-core: Return suppressed exception information in Closable.sequence ``PHAB_ID=D809561`` + +21.12.0 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: `Activity.collect*` and `Var.collect*` are now implemented in terms of known collection + type `scala.collection.Seq` versus HKT `CC[X]` before. This allows for certain performance + enhancements as well as makes it more aligned with the `Future.collect` APIs. + ``PHAB_ID=D795123`` + +21.11.0 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-security: Use snakeyaml to parse yaml instead of a buggy custom yaml + parser. This means that thrown IOExceptions have been replaced by + YAMLExceptions. Additionally, the parser member has been limited to private visibility. ``PHAB_ID=D617641`` + +New Features +~~~~~~~~~~~~ + +* util-security: Any valid yaml / json file with string keys and values can + be loaded with `com.twitter.util.security.Credentials`. ``PHAB_ID=D617641`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-cache: Update Caffeine cache library to version 2.9.2 ``PHAB_ID=D771893`` + +* util-jackson: Enable `BLOCK_UNSAFE_POLYMORPHIC_BASE_TYPES` in ScalaObjectMapper to + guard against Remote Code Execution (RCE) security vulnerability. This blocks + polymorphic deserialization from unsafe base types. ``PHAB_ID=D780863`` + +21.10.0 +------- + +New Features +~~~~~~~~~~~~ + +* util-core: Add convenience methods to convert to java.time.ZonedDateTime and + java.time.OffsetDateTime. `toInstant`, `toZonedDateTime`, and `toOffsetDateTime` also preserve + nanosecond resolution. ``PHAB_ID=D757636`` + +* util-stats: Moved `c.t.finagle.stats.LoadedStatsReceiver` and `c.t.finagle.stats.DefaultStatsReceiver` + from the finagle-core module to util-stats. ``PHAB_ID=D763497`` + +* util-jvm: Update exported metric schemas to properly reflect their underlying counter types + ``PHAB_ID=D710019`` + +21.9.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-jvm: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D728602`` + +* util-stats: Counter, Gauge, and Stat can be instrumented with descriptions. ``PHAB_ID = D615481`` + +* util-cache: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D714304``. + +* util-cache-guava: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D716101``. + +* util-routing: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D727120`` + +* util-sl4j-api: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D725491`` + +* util-sl4j-jul-bridge: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D726468`` + +* util-stats: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D720514`` + +* util-zk-test: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D720603`` + +* util-app: Flags parsing will now roll-up multiple flag parsing errors into a single + error message. When an error is encountered, flag parsing will continue to collect parse error + information instead of escaping on the first flag failure. After parsing all flags, if any errors + are present, a message containing all of the failed flags and their error reason, + along with the `help` usage message will be emitted. ```PHAB_ID=D729700`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.11.4. ``PHAB_ID=D727879`` + +* util: Bump version of json4s to 3.6.11. ``PHAB_ID=D729764`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-app: the `c.t.app.App#flags` field is now a `final def` instead of a `val` to address + `override val` scenarios where the `c.t.app.App#flags` are accessed as part of construction, + resulting in a NullPointerException due to access ordering issues. + The `c.t.app.Flags` class is now made final. The `c.t.app.App#name` field has changed from + a `val` and is now a `def`. A new `c.t.app.App#includeGlobalFlags` def has been exposed, which + defaults to `true`. The `c.t.app#includeGlobalFlags` def can be overridden to `false` + (ex: `override protected def includeGlobalFlags: Boolean = false`) in order to skip discovery + of `GlobalFlag`s during flag parsing. ```PHAB_ID=D723956`` + +21.8.0 (No 21.7.0 Release) +-------------------------- + +New Features +~~~~~~~~~~~~ + +* util-stats: add getAllExpressionsWithLabel utility to InMemoryStatsReceiver. +``PHAB_ID=D722942`` + +* util-app: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D714780`` + +* util-app-lifecycle: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D716444`` + +* util-codec: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D715114`` + +* util-hashing: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D718914`` + +* util-lint: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D698954`` + +* util-registry: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D716019`` + +* util-thrift: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D715129`` + +* util-app: Introduce a new `Command` class which provides a `Reader` interface to the output + of a shell command. ``PHAB_ID=D686134`` + +* util-core: Experimentally crossbuilds with Scala 3. ``PHAB_ID=D694775`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-app: Flags and GlobalFlag now use ClassTag instead of Manifest. ``PHAB_ID=D714780`` + +* util-thrift: ThriftCodec now uses ClassTag instead of Manifest. In + scala3 Manifest is intended for use by the compiler and should not be used in + client code. ``PHAB_ID=D715129`` + +* util-core (BREAKING): Remove `AbstractSpool`. Java users should use `Spools` static class or + the Spool companion object to create instances of `Spool`. ``PHAB_ID=D694775`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Update ScalaCheck to version 1.15.4 ``PHAB_ID=D691691`` + +* util-jackson: `JsonDiff#toSortedString` now includes null-type nodes, so that + `JsonDiff.Result#toString` shows differences in objects due to such nodes. ``PHAB_ID=D707033`` + +21.6.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-core: Add `ClasspathResource`, a utility for loading classpath resources as an + optional `InputStream`. ``PHAB_ID=D687324`` + +* util-jackson: Add `com.twitter.util.jackson.YAML` for YAML serde operations with a + default configured `ScalaObjectMapper.` Add more methods to `com.twitter.util.jackson.JSON` + ``PHAB_ID=D687327`` + +* util-jackson: Introduce a new library for JSON serialization and deserialization based on the + Jackson integration in `Finatra `__. + + This includes a custom case class deserializer which "fails slow" to collect all mapping failures + for error reporting. This deserializer is also natively integrated with the util-validator library + to provide for performing case class validations during deserialization. ``PHAB_ID=D664962`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-stats: Removed MetricSchema trait (CounterSchema, GaugeSchema and HistogramSchema). + StatReceiver derived classes use MetricBuilder directly to create counters, gauges and stats. + ``PHAB_ID=D668739`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-cache: Update Caffeine cache library to version 2.9.1 ``PHAB_ID=D660908`` + +21.5.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-validator: Introduce new library for case class validations (akin to Java bean validation) + which follows the Jakarta Bean Validation specification (https://beanvalidation.org/) by wrapping + the Hibernate Validator library and thus supports `jakarta.validation.Constraint` annotations and + validators for annotating and validating fields of Scala case classes. ``PHAB_ID=D638603`` + +* util-app: Introduce a Java-friendly API `c.t.app.App#runOnExit(Runnable)` and + `c.t.app.App#runOnExitLast(Runnable)` to allow Java 8 users to call `c.t.app.App#runOnExit` + and `c.t.app.App#runOnExitLast` with lambda expressions. ``PHAB_ID=D511536`` + +21.4.0 +------ + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-reflect: Memoize `c.t.util.reflect.Types#isCaseClass` computation. ``PHAB_ID=D657748`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-stats: Added a methods `c.t.f.stats.Counter#metadata: Metadata`, + `c.t.f.stats.Stat#metadata: Metadata`, and `c.t.f.stats.Gauge#metadata: + Metadata` to make it easier to introspect the constructed metric. In + particular, this will enable constructing `Expression`s based on the full name + of the metric. If you don't have access to a concrete `Metadata` instance + (like `MetricBuilder`) for constructing a Counter, Stat, or Gauge, you can + instead supply `NoMetadata`. ``PHAB_ID=D653751`` + +New Features +~~~~~~~~~~~~ + +* util-stats: Added a `com.twitter.finagle.stats.Metadata` abstraction, that can + be either many `com.twitter.finagle.stats.Metadata`, a `MetricBuilder`, or a + `NoMetadata`, which is the null `Metadata`. This enabled constructing + metadata for counters that represent multiple counters under the hood. + ``PHAB_ID=D653751`` + +21.3.0 +------ + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Revert to scala version 2.12.12 due to https://github.com/scoverage/sbt-scoverage/issues/319 + ``PHAB_ID=D635917`` + +* util: Bump scala version to 2.12.13 ``PHAB_ID=D632567`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util: Rename `c.t.util.reflect.Annotations#annotationEquals` to `c.t.util.reflect.Annotations#equals` + and `c.t.util.reflect.Types.eq` to `c.t.util.reflect.Types.equals`. ``PHAB_ID=D640386`` + +* util: Builds are now only supported for Scala 2.12+ ``PHAB_ID=D631091`` + +* util-reflect: Remove deprecated `c.t.util.reflect.Proxy`. There is no library replacement. + ``PHAB_ID=D630143`` + +* util-security: Renamed `com.twitter.util.security.PemFile` to `c.t.u.security.PemBytes`, and + changed its constructor to accept a string and a name. The main change here is that we assume + the PEM-encoded text has been fully buffered. To migrate, please use the helper method on the + companion object, `PemBytes#fromFile`. Note that unlike before with construction, we read from + the file, so it's possible for it to throw. ``PHAB_ID=D641088`` + +* util-security: Make KeyManager, TrustManager Factory helper methods public. The main change here + is that we expose util methods that are used to generate TrustManagerFactory, KeyManagerFactory + inside `c.t.u.security` for public usage. ``PHAB_ID=D644502`` + +New Features +~~~~~~~~~~~~ + +* util-reflect: Add `c.t.util.reflect.Annotations` a utility for finding annotations on a class and + `c.t.util.reflect.Classes` which has a utility for obtaining the `simpleName` of a given class + across JDK versions and while handling mangled names (those with non-supported Java identifier + characters). Also add utilities to determine if a given class is a case class in + `c.t.util.reflect.Types`. ``PHAB_ID=D638655`` + +* util-reflect: Add `c.t.util.reflect.Types`, a utility for some limited reflection based + operations. ``PHAB_ID=D631819`` + +* util-core: `c.t.io` now supports creating and deconstructing unsigned 128-bit buffers + in Buf. ``PHAB_ID=D606905`` + +* util-core: `c.t.io.ProxyByteReader` and `c.t.io.ProxyByteWriter` are now public. They are + useful for wrapping an existing `ByteReader` or `ByteWriter` and extending its functionality + without modifying the underlying instance. ``PHAB_ID=D622705`` + +* util-core: Provided `c.t.u.security.X509CertificateDeserializer` to make it possible to directly + deserialize an `X509Certificate` even if you don't have a file on disk. Also provided + `c.t.u.security.X509TrustManagerFactory#buildTrustManager` to make it possible to directly + construct an `X509TrustManager` with an `X509Certificate` instead of passing in a `File`. + ``PHAB_ID=D641088`` + +21.2.0 +------ + +No Changes + +21.1.0 +------ + +No Changes + +20.12.0 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ +* util-core: removed `com.twitter.util.Config.` ``PHAB_ID=D580444`` + +* util-core: removed `com.twitter.util.Future.isDone` method. The semantics of this method + are surprising in that `Future.exception(throwable).isDone == false`. Replace usages with + `Future.isDefined` or `Future.poll` ``PHAB_ID=D585700`` + +* util-stats: Changed `com.twitter.finagle.stats.MethodBuilder#withName`, + `com.twitter.finagle.stats.MethodBuilder#withRelativeName`, + `com.twitter.finagle.stats.MethodBuilder#withPercentiles`, + `com.twitter.finagle.stats.MethodBuilder#counter`, and + `com.twitter.finagle.stats.MethodBuilder#histogram`, to accept varargs as parameters, + rather than a `Seq`. Added a Java-friendly `com.twitter.finagle.stats.MethodBuilder#gauge`. + ``PHAB_ID=D620425`` + +New Features +~~~~~~~~~~~~ + +* util-core: `c.t.conversions` now includes conversion methods for maps (under `MapOps`) + that were moved from Finatra. ``PHAB_ID=D578819`` + +* util-core: `c.t.conversions` now includes conversion methods for tuples (under `TupleOps`) + that were moved from Finatra. ``PHAB_ID=D578804`` + +* util-core: `c.t.conversions` now includes conversion methods for seqs (under `SeqOps`) + that were moved from Finatra. ``PHAB_ID=D578605`` + +* util-core: `c.t.conversions` now includes conversion methods `toOption`, and `getOrElse` + under `StringOps`. ``PHAB_ID=D578549`` + +* util-core: `c.t.util.Duration` now includes `fromJava` and `asJava` conversions to + `java.time.Duration` types. ``PHAB_ID=D571885`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: `Activity.apply(Event)` will now propagate registry events to the underlying + `Event` instead of registering once and deregistering on garbage collection. This means + that if the underlying `Event` is "notified" while the derived `Activity` is not actively + being observed, it will not pick up the notification. Furthermore, the derived `Activity` + will revert to the `Activity.Pending` state while it is not under observation. ``PHAB_ID=D574843`` + +* util-core: `Activity#stabilize` will now propagate registry events to the underlying + `Activity` instead of registering once and deregistering on garbage collection. This means + that if the underlying `Activity` is changed to a new state while the derived `Activity` is not actively + being observed, it will not update its own state. The derived `Activity` will maintain its last + "stable" state when it's next observed, unless the underlying `Activity` was updated to a new "stable" + state, in which case it will pick that up instead. ``PHAB_ID=D574843`` + +* util-stats: `c.t.finagle.stats.DenylistStatsReceiver` now includes methods for creating + `DenyListStatsReceiver` from partial functions. ``PHAB_ID=D576833`` + +* util-core: `c.t.util.FuturePool` now supports exporting the number of its pending tasks via + `numPendingTasks`. ``PHAB_ID=D583030`` + +20.10.0 +------- + +Bug Fixes +~~~~~~~~~ +* util-stat: `MetricBuilder` now uses a configurable metadataScopeSeparator to align + more closely with the metrics.json api. Services with an overridden scopeSeparator will + now see that reflected in metric_metadata.json where previously it was erroneously using + / in all cases. ``PHAB_ID=D557994`` + +* util-slf4j-api: Better Java interop. Deprecate `c.t.util.logging.Loggers` as Java users should be + able to use the `c.t.util.logging.Logger` companion object with less verbosity required. + ``PHAB_ID=D558605`` + +20.9.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-app: `Seq`/`Tuple2`/`Map` flags can now operate on booleans. For example, + `Flag[Seq[Boolean]]` now works as expected instead of throwing an assert exception (previous + behaviour). ``PHAB_ID=D549196`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-app: `Flaggable.mandatory` now takes implicit `ClassTag[T]` as an argument. This is change is + source-compatible in Scala but requires Java users to pass argument explicitly via + `ClassTag$.MODULE$.apply(clazz)`. ``PHAB_ID=D542135`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump version of Jackson to 2.11.2. ``PHAB_ID=D538440`` + +20.8.1 +------ + +New Features +~~~~~~~~~~~~ + +* util-mock: Introduce `mockito-scala `__ based mocking + integration. Fix up and update mockito testing dependencies: + - `mockito-all:1.10.19` to `mockito-core:3.3.3` + - `scalatestplus:mockito-1-10:3.1.0.0` to `scalatestplus:mockito-3-2:3.1.2.0` ``PHAB_ID=D530995`` + +* util-app: Add support for flags of Java Enum types. ``PHAB_ID=D530205`` + +20.8.0 (DO NOT USE) +------------------- + +New Features +~~~~~~~~~~~~ +* util-stats: Store MetricSchemas in InMemoryStatsReceiver to enable further testing. ``PHAB_ID=D518195`` + +* util-core: c.t.u.Var.Observer is now public. This allows scala users to extend the Var trait + as has been the case for Java users. ``PHAB_ID=D520237`` + +* util-core: Added two new methods to c.t.u.Duration and c.t.u.Time: `.fromHours` and `.fromDays`. ``PHAB_ID=D522734`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ +* util-app: Treat empty strings as empty collections in `Flag[Seq[_]]`, `Flag[Set[_]]`, + `Flag[java.util.List[_]]`, and `Flag[java.util.Set[_]]`. They were treated as collections + with single element, empty string, before. ``PHAB_ID=D516724`` + +20.7.0 +------ + +New Features +~~~~~~~~~~~~ +* util-app: Add ability to observe App lifecycle events. ``PHAB_ID=D508145`` + +20.6.0 +------ + +New Features +~~~~~~~~~~~~ +* util-stats: Add two new Java-friendly methods to `StatsReceiver` (`addGauge` and `provideGauge`) + that take java.util.function.Supplier as well as list vararg argument last to enable better + developers' experience. ``PHAB_ID=D497885`` + +* util-app: Add a `Flaggable` instance for `java.time.LocalTime`. ``PHAB_ID=D499606`` + +* util-app: Add two new methods to retrieve flag's unparsed value (as string): `Flag.getUnparsed` + and `Flag.getWithDefaultUnparsed`. ``PHAB_ID=D499628`` + +20.5.0 +------ + +New Features +~~~~~~~~~~~~ +* util-security: Moved Credentials from util-core + `c.t.util.Credentials` => `c.t.util.security.Credentials`. + ``PHAB_ID=D477984`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ +* util-core: Move Credentials to util-security + `c.t.util.Credentials` => `c.t.util.security.Credentials`. + ``PHAB_ID=D477984`` + +* util-core: Change the namespace of `ActivitySource` and its derivatives to + `com.twitter.io` as its no longer considered experimental since the code has + changed minimally in the past 5 years. ``PHAB_ID=D478498`` + +20.4.1 +------ + +New Features +~~~~~~~~~~~~ +* util-tunable: ConfigurationLinter accepts a relative path. ``PHAB_ID=D468528`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Bump jackson version to 2.11.0. ``PHAB_ID=D457496`` + +20.4.0 (DO NOT USE) +------------------- + +New Features +~~~~~~~~~~~~ +* util-core: When looking to add idempotent close behavior, users should mix in `CloseOnce` to + classes which already extend (or implement) `Closable`, as mixing in `ClosableOnce` leads to + compiler errors. `ClosableOnce.of` can still be used to create a `ClosableOnce` proxy of an + already instantiated `Closable`. Classes which do not extend `Closable` can still + mix in `ClosableOnce`. ``PHAB_ID=D455819`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ +* util-hashing: Rename + `c.t.hashing.KetamaNode` => `HashNode`, + `c.t.hashing.KetamaDistributor` => `ConsistentHashingDistributor`. + ``PHAB_ID=D449929`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-stats: Provide CachedRegex, a function that filters a + Map[String, Number] => Map[String, Number] and caches which keys to filter on + based on a regex. Useful for filtering down metrics in the style that Finagle + typically recommends ``PHAB_ID=D459391``. + +20.3.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: The system property `com.twitter.util.UseLocalInInterruptible` no longer + can be used to modulate which Local state is present when a Promise is interrupted. + ``PHAB_ID=D442444`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: Promises now exclusively use the state local to setting the interrupt + handler when raising on a Promise. ``PHAB_ID=D442444`` + +20.2.1 +------ + +New Features +~~~~~~~~~~~~ + +* util-app: Add `c.t.util.app.App#onExitLast` to be able to provide better Java + ergonomics for designating a final exit function. ``PHAB_ID=D433874`` + +* util-core: Add `c.t.io.Reader.concat` to conveniently concatenate a collection + of Reader to a single Reader. ``PHAB_ID=D434448`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: `Future.unapply` has been removed. Use `Future.poll` to retrieve Future's + state. ``PHAB_ID=D427429`` + +Bug Fixes +~~~~~~~~~ + +* util-logging: Add a missing `_*` that could result in exceptions when using the + `Logger.apply(Level, Throwable, String, Any*)` signature. ``PHAB_ID=D430122`` + +20.1.0 +------ + +No Changes + +19.12.0 +------- + +New Features +~~~~~~~~~~~~ + +* util-stats: Introduces `c.t.f.stats.LazyStatsReceiver` which ensures that counters and histograms + don't export metrics until after they have been `incr`ed or `add`ed at least once. ``PHAB_ID=D398898`` + +* util-core: Introduce `Time#nowNanoPrecision` to produce nanosecond resolution timestamps in JDK9 + or later. ``PHAB_ID=D400661`` + +* util-core: Introduce `Future#toCompletableFuture`, which derives a `CompletableFuture` from + a `com.twitter.util.Future` to make integrating with Java APIs simpler. ``PHAB_ID=D408656`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Upgrade to jackson 2.9.10 and jackson-databind 2.9.10.1 ``PHAB_ID=D410846`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: The lightly used `com.twitter.util.JavaSingleton` trait has been removed. It + did not work as intended. Users should provide Java friendly objects, classes, and methods + instead. ``PHAB_ID=D399947`` + +Deprecations +~~~~~~~~~~~~ + +* util-test: The `c.t.logging.TestLogging` mixin has been deprecated. Users are encouraged to + move to slf4j for logging and minimize dependencies on `com.twitter.logging` in general, as + it is intended to be replaced entirely by slf4j. ``PHAB_ID=D403574`` + +Bug Fixes +~~~~~~~~~ + +* util-core: `Future#toJavaFuture` incorrectly threw the exception responsible for failing it, + instead of a `j.u.c.ExecutionException` wrapping the exception responsible for failing it. + ``PHAB_ID=D408656`` + +19.11.0 +------- + +New Features +~~~~~~~~~~~~ + +* util: Add initial support for JDK 11 compatibility. ``PHAB_ID=D365075`` + +* util-core: Created public method Closable.stopCollectClosablesThread that stops CollectClosables +thread. ``PHAB_ID=D382800`` + +* util-core: Introduced `Reader.fromIterator` to create a Reader from an iterator. It is not +recommended to call `iterator.next()` after creating a `Reader` from it. Doing so will affect the +behavior of `Reader.read()` because it will skip the value returned from `iterator.next`. +``PHAB_ID=D391769`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Upgrade to caffeine 2.8.0 ``PHAB_ID=D384592`` + +* util-jvm: Stop double-exporting `postGC` stats under both `jvm` and `jvm/mem`. These are now + only exported under `jvm/mem/postGC`. ``PHAB_ID=D392230`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-stats: abstract methods of StatsReceiver now take Schemas. The old APIs + are now final and cannot be overriden. For custom implementations, define + schema based methods (eg, counter(verbosity, name) is now + counter(CounterSchema)). NB: users can continue to call the old interface; + only implementors must migrate.``PHAB_ID=D385068`` + +* util-core: Removed `c.t.io.Pipe.copyMany` (was `Reader.copyMany`). Use `AsyncStream.foreachF` + link to `Pipe.copy` for substitution. ``PHAB_ID=D396590`` + +* util-core: Add `c.t.io.BufReader.readAll` to consume a `Reader[Buf]` and concat values to a Buf. + Replace `c.t.io.Reader.readAll` with `Reader.readAllItems`, the new API consumes a generic Reader[T], + and return a Seq of items. ``PHAB_ID=D391346`` + +* util-core: Moved `c.t.io.Reader.chunked` to `c.t.io.BufReader.chunked`, and `Reader.framed` to + `BufReader.framed`. ``PHAB_ID=D392198`` + +* util-core: Moved `c.t.io.Reader.copy` to `c.t.io.Pipe.copy`, and `Reader.copyMany` to + `Pipe.copyMany`. ``PHAB_ID=D393650`` + +Deprecations +~~~~~~~~~~~~ + +* util-core: Mark `c.t.io.BufReaders`, `c.t.io.Bufs`, `c.t.io.Readers`, and `c.t.io.Writers` as + Deprecated. These classes will no longer be needed, and will be removed, after 2.11 support is + dropped. ``PHAB_ID=D393913`` + +* util-stats: Removed deprecated methods `stat0` and `counter0` from `StatsReceiver`. ``PHAB_ID=D393063`` + +19.10.0 +------- + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: When a computation from FuturePool is interrupted, its promise is + set to the interrupt, wrapped in a j.u.c.CancellationException. This wrapper + was introduced because, all interrupts were once CancellationExceptions. In + RB_ID=98612, this changed to allow the user to raise specific exceptions as + interrupts, and in the aid of compatibility, we wrapped this raised exception + in a CancellationException. This change removes the wrapper and fails the + promise directly with the raised exception. This will affect users that + explicitly handle CancellationException. ``PHAB_ID=D371872`` + +Bug Fixes +~~~~~~~~~ + +* util-core: Fixed bug in `c.t.io.Reader.framed` where if the `framer` didn't emit a `List` the + emitted frames were skipped. ``PHAB_ID=D378048`` + +* util-hashing: Fix a bug where `partitionIdForHash` was returning incosistent values w.r.t + `entryForHash` in `KetamaDistributor`. ``PHAB_ID=D381128`` + +19.9.0 +------ + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-app: Better handling of exceptions when awaiting on the `c.t.app.App` to close at + the end of the main function. We `Await.ready` on `this` as the last step of + `App#nonExitingMain` which can potentially throw a `TimeoutException` which was previously + unhandled. We have updated the logic to ensure that `TimeoutException`s are handled accordingly. + ``PHAB_ID=D356846`` + +* util: Upgrade to Scala Collections Compat 2.1.2. ``PHAB_ID=D364013`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: BoundedStack is unused and really old code. Delete it. ``PHAB_ID=D357338`` + +* util-logging: `com.twitter.logging.ScribeHandler` and `com.twitter.logging.ScribeHandlers` have + been removed. Users are encouraged to use slf4j for logging. However, if a util-logging integrated + ScribeHandler is still required, users can either build their own Finagle-based scribe client as + in `ScribeRawZipkinTracer` in finagle-zipkin-scribe, or copy the old `ScribeHandler` + implementation directly into their code. ``PHAB_ID=D357008`` + +19.8.1 +------ + +New Features +~~~~~~~~~~~~ + +* util: Enables cross-build for 2.13.0. ``PHAB_ID=D333021`` + +Java Compatibility +~~~~~~~~~~~~~~~~~~ + +* util-stats: In `c.t.finagle.stats.AbstractStatsReceiver`, the `counter`, `stat` and + `addGauge` become final, override `counterImpl`, `statImpl` and `addGaugeImpl` instead. + ``PHAB_ID=D333021`` + +* util-core: + `c.t.concurrent.Offer.choose`, + `c.t.concurrent.AsyncStream.apply`, + `c.t.util.Await.all`, + `c.t.util.Closable.sequence` become available to java for passing varargs. ``PHAB_ID=D333021`` + +* util-stats: + `c.t.finagle.stats.StatsReceiver.provideGauge` and `addGauge` become available to java for + passing varags. ``PHAB_ID=D333021`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: (not breaking) `c.t.util.Future.join` and `c.t.util.Future.collect` now take + `Iterable[Future[A]]` other than Seq. ``PHAB_ID=D333021`` + +* util-core: Revert the change above, in `c.t.util.Future`, `collect`, `collectToTry` and `join` + take `scala.collection.Seq[Future[A]]`. ``PHAB_ID=D355403`` + +* util-core: `com.twitter.util.Event#build` now builds a Seq of events. `Event#buildAny` builds + against any collection of events. ``PHAB_ID=D333021`` + +19.8.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-logging: The namespace forwarders for `Level` and `Policy` in `com.twitter.logging.config` + have been removed. Code should be updated to use `com.twitter.logging.Level` and + `com.twitter.logging.Policy` where necessary. Users are encouraged to use 'util-slf4j-api' though + where possible. ``PHAB_ID=D344439`` + +* util-logging: The deprecated `com.twitter.logging.config.LoggerConfig` and associated + classes have been removed. These have been deprecated since 2012. Code should be updated + to use `com.twitter.logging.LoggerFactory` where necessary. Users are encouraged to use + 'util-slf4j-api' though where possible. ``PHAB_ID=D345381`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util: Upgrade to Jackson 2.9.9. ``PHAB_ID=D345969`` + +* util-app: It is now illegal to define GlobalFlags enclosed in package objects. ``PHAB_ID=D353045`` + +19.7.0 +------ + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: Removed deprecated `c.t.concurrent.Scheduler` methods `usrTime`, + `cpuTime`, and `wallTime`. These were deprecated in 2015 and have no + replacement. ``PHAB_ID=D330386`` + +* util-core: Removed deprecated `com.twitter.logging.config` classes `SyslogFormatterConfig`, + `ThrottledHandlerConfig`, `SyslogHandlerConfig`. These were deprecated in 2012 and have + no replacement. Users are encouraged to use 'util-slf4j-api' where possible. ``PHAB_ID=D339563`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: Remove experimental toggle `com.twitter.util.BypassScheduler` used + for speeding up `ConstFuture.map` (`transformTry`). Now, we always run map + operations immediately instead of via the Scheduler, where they may be queued + and potentially reordered. ``PHAB_ID=D338487`` + +19.6.0 +------ + +Bug Fixes +~~~~~~~~~ + +* util-core: Fixed the behavior in `c.t.io.Reader` where reading from `Reader#empty` fails to return + a `ReaderDiscardedException` after it's discarded. ``PHAB_ID=D325465`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: Use Local at callback creation for Future's interrupt handler rather than + raiser's locals so that it is consistent with other callbacks. This functionality is + currently disabled and can be enabled by a toggle (com.twitter.util.UseLocalInInterruptible) + by setting it to 1.0 if you would like to try it out. ``PHAB_ID=D324315`` + +19.5.1 +------ + +No Changes + +19.5.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-app: Track the registration of duplicated Flag names. Currently, we print a warning to + `stderr` but do not track the duplicated Flag names. Tracking them allows us to inspect and + warn over the entire set. ``PHAB_ID=D314410`` + +19.4.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-app: Improve usage of `Flag.let` by providing a `Flag.letParse` method + ``PHAB_ID=D288549`` + +19.3.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-core: Discard parent reader from `Reader.flatten` when child reader encounters an exception. + ``PHAB_ID=D281830`` + +* util-core: Added `c.t.conversions.StringOps#toSnakeCase,toCamelCase,toPascalCase` implementations. + ``PHAB_ID=D280886`` + +19.2.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-core: updated `Reader#fromFuture` to resolve its `onClose` when reading of end-of-stream. + ``PHAB_ID=D269413`` + +* util-core: Added `Reader.flatten` to flatten a `Reader[Reader[_]]` to `Reader[_]`, + and `Reader.fromSeq` to create a new Reader from a Seq. ``PHAB_ID=D255424`` + +* util-core: Added `Duration.fromMinutes` to return a `Duration` from a given number of minutes. + ``PHAB_ID=D259795`` + +* util-core: If given a `Timer` upon construction, `c.t.io.Pipe` will respect the close + deadline and wait the given amount of time for any pending writes to be read. ``PHAB_ID=D229728`` + +* util-core: Optimized `ConstFuture.proxyTo` which brings the performance of + `flatMap` and `transform` of a `ConstFuture` in line with `map`. ``PHAB_ID=D271358`` + +* util-core: Experimental toggle (com.twitter.util.BypassScheduler) for speeding up + `ConstFuture.map` (`transformTry`). The mechanism, when turned on, runs map operations + immediately (why not when we have a concrete value), instead of via the Scheduler, where it may + be queued and potentially reordered, e.g.: + `f.flatMap { _ => println(1); g.map { _ => println(2) }; println(3) }` will print `1 2 3`, + where it would have printed `1 3 2`. ``PHAB_ID=D271962`` + +* util-security: `Pkcs8KeyManagerFactory` now supports a certificates file which contains multiple + certificates that are part of the same certificate chain. ``PHAB_ID=D263190`` + +Bug Fixes +~~~~~~~~~ + +* util-core: Fixed the behavior in `c.t.io.Reader` where `Reader#flatMap` fails to propagate + parent reader's `onClose`. ``PHAB_ID=D269413`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: Closing a `c.t.io.Pipe` will notify `onClose` when the deadline has passed whereas + before the pipe would wait indefinitely for a read before transitioning to the Closed state. + ``PHAB_ID=D229728`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: Remove `c.t.u.CountDownLatch` which is an extremely thin shim around + `j.u.c.CountDownLatch` that provides pretty limited value. To migrate to `j.u.c.CountDownLatch`, + instead of `c.t.u.CountDownLatch#await(Duration)`, please use + `j.u.c.CountDownLatch#await(int, TimeUnit)`, and instead of + `c.t.u.CountDownLatch#within(Duration)`, please throw an exception yourself after awaiting. + ``PHAB_ID=D269404`` + +* util-core: Deprecated conversions in `c.t.conversions` have new implementations + that follow a naming scheme of `SomethingOps`. ``PHAB_ID=D272206`` + + - `percent` is now `PercentOps` + - `storage` is now `StorageUnitOps` + - `string` is now `StringOps` + - `thread` is now `ThreadOps` + - `time` is now `DurationOps` + - `u64` is now `U64Ops` + +* util-collection: Delete util-collection. We deleted `GenerationalQueue`, `MapToSetAdapter`, and + `ImmutableLRU`, because we found that they were of little utility. We deleted `LruMap` because it + was a very thin shim around a `j.u.LinkedHashMap`, where you override `removeEldestEntry`. If you + need `SynchronizedLruMap`, you can wrap your `LinkedHashMap` with + `j.u.Collection.synchronizedMap`. We moved `RecordSchema` into finagle-base-http because it was + basically only used for HTTP messages, so its new package name is `c.t.f.http.collection`. + ``PHAB_ID=D270548`` + +* util-core: Rename `BlacklistStatsReceiver` to `DenylistStatsReceiver`. ``PHAB_ID=D270526`` + +* util-core: `Buf.Composite` is now private. Program against more generic, `Buf` interface instead. + ``PHAB_ID=D270916`` + +19.1.0 +------ + +New Features +~~~~~~~~~~~~ + +* util-core: Added Reader.map/flatMap to transform Reader[A] to Reader[B]. Added `fromFuture()` + and `value()` in the Reader object to construct a new Reader. ``PHAB_ID=D252165`` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: The implicit conversions classes in `c.t.conversions.SomethingOps` have been + renamed to have unique names. This allows them to be used together with wildcard imports. + See Github issue (https://github.com/twitter/util/issues/239). ``PHAB_ID=D252462`` + +* util-core: Both `c.t.io.Writer.FailingWriter` and `c.t.io.Writer.fail` were removed. Build your + own instance should you need to. ``PHAB_ID=D256615`` + +18.12.0 +------- + +New Features +~~~~~~~~~~~~ + +* util-core: Provide a way to listen for stream termination to `c.t.util.Reader`, `Reader#onClose` + which is satisfied when the stream is discarded or read until the end. ``PHAB_ID=D236311`` + +* util-core: Conversions in `c.t.conversions` have new implementations + that follow a naming scheme of `SomethingOps`. Where possible the implementations + are `AnyVal` based avoiding allocations for the common usage pattern. + ``PHAB_ID=D249403`` + + - `percent` is now `PercentOps` + - `storage` is now `StorageUnitOps` + - `string` is now `StringOps` + - `thread` is now `ThreadOps` + - `time` is now `DurationOps` + - `u64` is now `U64Ops` + +Bug Fixes +~~~~~~~~~ + +* util-core: Fixed a bug where tail would sometimes return Some empty AsyncStream instead of None. + ``PHAB_ID=D241513`` + +Deprecations +~~~~~~~~~~~~ + +* util-core: Conversions in `c.t.conversions` have been deprecated in favor of `SomethingOps` + versions. Where possible the implementations are `AnyVal` based and use implicit classes + instead of implicit conversions. ``PHAB_ID=D249403`` + + - `percent` is now `PercentOps` + - `storage` is now `StorageUnitOps` + - `string` is now `StringOps` + - `thread` is now `ThreadOps` + - `time` is now `DurationOps` + - `u64` is now `U64Ops` + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: Experimental `c.t.io.exp.MinThroughput` utilities were removed. ``PHAB_ID=D240944`` + +* util-core: Deleted `c.t.io.Reader.Null`, which was incompatible with `Reader#onClose` semantics. + `c.t.io.Reader#empty[Nothing]` is a drop-in replacement. ``PHAB_ID=D236311`` + +* util-core: Removed `c.t.util.U64` bits. Use `c.t.converters.u64._` instead. ``PHAB_ID=D244723`` + +18.11.0 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: `c.t.u.Future.raiseWithin` methods now take the timeout exception as a call-by-name + parameter instead of a strict exception. While Scala programs should compile as usual, Java + users will need to use a `scala.Function0` as the second parameter. The helper + `c.t.u.Function.func0` can be helpful. ``PHAB_ID=D229559`` + +* util-core: Rename `c.t.io.Reader.ReaderDiscarded` to `c.t.io.ReaderDiscardedException`. + ``PHAB_ID=D231969`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: Made Stopwatch.timeNanos monotone. ``PHAB_ID=D236629`` + +18.10.0 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: `c.t.io.Reader.Writable` and `c.t.Reader.writable()` are removed. Use `c.t.io.Pipe` + instead. ``PHAB_ID=D226603`` + +* util-core: `c.t.util.TempFolder` has been moved to `c.t.io.TempFolder`. ``PHAB_ID=D226940`` + +* util-core: Removed the forwarding types `c.t.util.TimeConversions` and + `c.t.util.StorageUnitConversions`. Use `c.t.conversions.time` and + `c.t.conversions.storage` directly. ``PHAB_ID=D227363`` + +* util-core: `c.t.concurrent.AsyncStream.fromReader` has been moved to + `c.t.io.Reader.toAsyncStream`. ``PHAB_ID=D228277`` + +* util-core: `c.t.io.Reader.read()` no longer takes `n`, the maximum number of bytes to read off a + stream. ``PHAB_ID=D228385`` + +New Features +~~~~~~~~~~~~ + +* util-core: `c.t.io.Reader.fromBuf` (`BufReader`), `c.t.io.Reader.fromFile`, + `c.t.io.Reader.fromInputStream` (`InputStreamReader`) now take an additional parameter, + `chunkSize`, the upper bound of the number of bytes that a given reader emits at each read. + ``PHAB_ID=D203154`` + +Runtime Behavior Changes +~~~~~~~~~~~~~~~~~~~~~~~~ + +* util-core: `c.t.u.Duration.inTimeUnit` can now return + `j.u.c.TimeUnit.MINUTES`. ``PHAB_ID=D225115`` + +18.9.1 +------- + +Breaking API Changes +~~~~~~~~~~~~~~~~~~~~ + +* util-core: `c.t.io.Writer` now extends `c.t.util.Closable`. `c.t.io.Writer.ClosableWriter` + is no longer exist. ``PHAB_ID=D218453`` + +* util-core: Add `onClose` into `c.t.io.Writer`, it exposes a `Future` that is satisfied when + the stream is closed. ``PHAB_ID=D226319`` + +Bug Fixes +~~~~~~~~~ + +* util-slf4j-api: Moved slf4j-simple dependency to be a 'test' dependency, instead of a + compile dependency, which was inaccurate. ``PHAB_ID=D220718`` + +New Features +~~~~~~~~~~~~ + +* util-core: Added a `contramap` function into `c.t.io.Writer`, `Writer` is now a contravariant + functor. Added the `AbstractWriter` for Java compatibility ``PHAB_ID=D225686`` + 18.9.0 ------- @@ -28,6 +1227,9 @@ New Features * util-security: Add `NullSslSession` related objects for use with non-existent `SSLSession`s. ``PHAB_ID=D201421`` +* util-tunable: Introducing `Tunable.asVar` that allows observing changes to tunables. + ``PHAB_ID=D211622`` + Breaking API Changes ~~~~~~~~~~~~~~~~~~~~ diff --git a/CONFIG.ini b/CONFIG.ini index d9780e87f6..0b96e72684 100644 --- a/CONFIG.ini +++ b/CONFIG.ini @@ -1,5 +1,3 @@ -; See http://go/CONFIG.ini - ; owned by Core System Libraries [jira] project: CSL diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index ed1abc0822..8edfa2a93e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -2,10 +2,24 @@ We'd love to get patches from you! +## Building + +This library is built using [sbt][sbt]. When building please use the included +[`./sbt`](https://github.com/twitter/util/blob/develop/sbt) script which +provides a thin wrapper over [sbt][sbt] and correctly sets memory and other +settings. + If you have any questions or run into any problems, please create an issue here, chat with us in [gitter](https://gitter.im/twitter/finagle), or email the Finaglers [mailing list](https://groups.google.com/forum/#!forum/finaglers). +### 3rd party upgrades + +We upgrade the following 3rd party libraries/tools at least once every 3 months: +* [sbt](https://github.com/sbt/sbt) +* [Jackson](https://github.com/FasterXML/jackson) +* [Caffeine](https://github.com/ben-manes/caffeine) + ## Workflow The workflow that we support: @@ -65,7 +79,7 @@ locally). When appropriate, use [ScalaCheck][scalacheck] to write property-based tests for your code. This will often produce more thorough and effective inputs for your tests. We use ScalaTest's -[GeneratorDrivenPropertyChecks][gendrivenprop] as the entry point for +[ScalaCheckDrivenPropertyChecks][gendrivenprop] as the entry point for writing these tests. ## Compatibility @@ -171,14 +185,15 @@ metadata will be preserved. Please let us know if you have any questions about this process! [0]: https://github.com/twitter/util/pull/109 -[1]: https://github.com/twitter/util/blob/master/util-core/src/test/java/com/twitter/util/VarCompilationTest.java -[scaladoc]: http://docs.scala-lang.org/style/scaladoc.html -[es]: https://twitter.github.io/effectivescala/ -[sd]: http://docs.scala-lang.org/style/scaladoc.html -[funsuite]: http://www.scalatest.org/getting_started_with_fun_suite -[scalatest]: http://www.scalatest.org/ -[ssg]: http://docs.scala-lang.org/style/scaladoc.html -[travis-ci]: https://travis-ci.org/twitter/util +[1]: https://github.com/twitter/util/blob/release/util-core/src/test/java/com/twitter/util/VarCompilationTest.java [changes]: https://github.com/twitter/util/blob/develop/CHANGELOG.rst +[es]: https://twitter.github.io/effectivescala/ +[funsuite]: https://www.scalatest.org/getting_started_with_fun_suite +[gendrivenprop]: https://www.scalatest.org/user_guide/generator_driven_property_checks +[sbt]: https://www.scala-sbt.org/ [scalacheck]: https://www.scalacheck.org/ -[gendrivenprop]: http://www.scalatest.org/user_guide/generator_driven_property_checks +[scaladoc]: https://docs.scala-lang.org/style/scaladoc.html +[scalatest]: https://www.scalatest.org/ +[sd]: https://docs.scala-lang.org/style/scaladoc.html +[ssg]: https://docs.scala-lang.org/style/scaladoc.html +[travis-ci]: https://travis-ci.org/twitter/util diff --git a/GROUPS b/GROUPS deleted file mode 100644 index e57fe3a541..0000000000 --- a/GROUPS +++ /dev/null @@ -1 +0,0 @@ -csl diff --git a/OWNERS b/OWNERS deleted file mode 100644 index c313f6c1d5..0000000000 --- a/OWNERS +++ /dev/null @@ -1,5 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta diff --git a/PROJECT.yml b/PROJECT.yml new file mode 100644 index 0000000000..675c29685a --- /dev/null +++ b/PROJECT.yml @@ -0,0 +1,6 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - csl-team:ldap + - jillianc + - mbezoyan diff --git a/README.md b/README.md index cdd84070fa..5960d3cbe2 100644 --- a/README.md +++ b/README.md @@ -1,14 +1,14 @@ # Twitter Util -[![Build status](https://travis-ci.org/twitter/util.svg?branch=master)](https://travis-ci.org/twitter/util) -[![Codecov](https://codecov.io/gh/twitter/util/branch/develop/graph/badge.svg)](https://codecov.io/gh/twitter/util) +[![Build Status](https://github.com/twitter/util/workflows/continuous%20integration/badge.svg?branch=develop)](https://github.com/twitter/util/actions?query=workflow%3A%22continuous+integration%22+branch%3Adevelop) [![Project status](https://img.shields.io/badge/status-active-brightgreen.svg)](#status) [![Gitter](https://badges.gitter.im/twitter/finagle.svg)](https://gitter.im/twitter/finagle?utm_source=badge&utm_medium=badge&utm_campaign=pr-badge) [![Maven Central](https://maven-badges.herokuapp.com/maven-central/com.twitter/util-core_2.12/badge.svg)](https://maven-badges.herokuapp.com/maven-central/com.twitter/util-core_2.12) A bunch of idiomatic, small, general purpose tools. -See the Scaladoc [here](https://twitter.github.com/util). +See the Scaladoc [here](https://twitter.github.io/util/docs/#com.twitter.util.package) +or check out the [user guide](https://twitter.github.io/util). ## Status @@ -18,24 +18,28 @@ and is being actively developed and maintained. ## Releases [Releases](https://maven-badges.herokuapp.com/maven-central/com.twitter/util_2.12) -are done on an approximately monthly schedule. While [semver](http://semver.org/) +are done on an approximately monthly schedule. While [semver](https://semver.org/) is not followed, the [changelogs](CHANGELOG.rst) are detailed and include sections on public API breaks and changes in runtime behavior. ## Contributing -The `master` branch of this repository contains the latest stable release of +We feel that a welcoming community is important and we ask that you follow Twitter's +[Open Source Code of Conduct](https://github.com/twitter/.github/blob/main/code-of-conduct.md) +in all interactions with the community. + +The `release` branch of this repository contains the latest stable release of Util, and weekly snapshots are published to the `develop` branch. In general pull requests should be submitted against `develop`. See -[CONTRIBUTING.md](https://github.com/twitter/util/blob/master/CONTRIBUTING.md) +[CONTRIBUTING.md](https://github.com/twitter/util/blob/release/CONTRIBUTING.md) for more details about how to contribute. # Using in your project -An example SBT dependency string for the `util-collection` tools would look like this: +An example SBT dependency string for the `util-core` library would look like this: ```scala -val collUtils = "com.twitter" %% "util-collection" % "18.9.0" +val utilCore = "com.twitter" %% "util-core" % "24.5.0" ``` # Units @@ -43,7 +47,7 @@ val collUtils = "com.twitter" %% "util-collection" % "18.9.0" ## Time ```scala -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ val duration1 = 1.second val duration2 = 2.minutes @@ -53,7 +57,7 @@ duration1.inMillis // => 1000L ## Space ```scala -import com.twitter.conversions.storage._ +import com.twitter.conversions.StorageUnitOps._ val amount = 8.megabytes amount.inBytes // => 8388608L amount.inKilobytes // => 8192L @@ -64,7 +68,7 @@ amount.inKilobytes // => 8192L A Non-actor re-implementation of Scala Futures. ```scala -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, Future, Promise} val f = new Promise[Int] @@ -89,18 +93,24 @@ for { } ``` -# Collections - -## LruMap +## Future interrupts -The LruMap is an LRU with a maximum size passed in. If the map is full it expires items in FIFO order. Reading a value will move an item to the top of the stack. +Method `raise` on `Future` (`def raise(cause: Throwable)`) raises the interrupt described by +`cause` to the producer of this `Future`. Interrupt handlers are installed on a `Promise` +using `setInterruptHandler`, which takes a partial function: ```scala -import com.twitter.util.LruMap - -val map = new LruMap[String, String](15) // this is of type mutable.Map[String, String] +val p = new Promise[T] +p.setInterruptHandler { + case exc: MyException => + // deal with interrupt.. +} ``` +Interrupts differ in semantics from cancellation in important ways: there can only be one +interrupt handler per promise, and interrupts are only delivered if the promise is not yet +complete. + # Object Pool The pool order is FIFO. @@ -171,209 +181,7 @@ val distributor = new KetamaDistributor(nodes, 1 /* num reps */) distributor.nodeForHash("abc".##) // => client ``` -# Logging - -`util-logging` is a small wrapper around Java's built-in logging to make it more -Scala-friendly. - -## Using - -To access logging, you can usually just use: - -```scala -import com.twitter.logging.Logger -private val log = Logger.get(getClass) -``` - -This creates a `Logger` object that uses the current class or object's -package name as the logging node, so class "com.example.foo.Lamp" will log to -node `com.example.foo` (generally showing "foo" as the name in the logfile). -You can also get a logger explicitly by name: - -```scala -private val log = Logger.get("com.example.foo") -``` - -Logger objects wrap everything useful from `java.util.logging.Logger`, as well -as adding some convenience methods: - -```scala -// log a string with sprintf conversion: -log.info("Starting compaction on zone %d...", zoneId) - -try { - ... -} catch { - // log an exception backtrace with the message: - case e: IOException => - log.error(e, "I/O exception: %s", e.getMessage) -} -``` - -Each of the log levels (from "fatal" to "trace") has these two convenience -methods. You may also use `log` directly: - -```scala -import com.twitter.logging.Level -log(Level.DEBUG, "Logging %s at debug level.", name) -``` - -An advantage to using sprintf ("%s", etc) conversion, as opposed to: - -```scala -log(Level.DEBUG, s"Logging $name at debug level.") -``` - -is that Java & Scala perform string concatenation at runtime, even if nothing -will be logged because the log file isn't writing debug messages right now. -With `sprintf` parameters, the arguments are just bundled up and passed directly -to the logging level before formatting. If no log message would be written to -any file or device, then no formatting is done and the arguments are thrown -away. That makes it very inexpensive to include verbose debug logging which -can be turned off without recompiling and re-deploying. - -If you prefer, there are also variants that take lazily evaluated parameters, -and only evaluate them if logging is active at that level: - -```scala -log.ifDebug(s"Login from $name at $date.") -``` - -The logging classes are done as an extension to the `java.util.logging` API, -and so you can use the Java interface directly, if you want to. Each of the -Java classes (Logger, Handler, Formatter) is just wrapped by a Scala class. - - -## Configuring - -In the Java style, log nodes are in a tree, with the root node being "" (the -empty string). If a node has a filter level set, only log messages of that -priority or higher are passed up to the parent. Handlers are attached to nodes -for sending log messages to files or logging services, and may have formatters -attached to them. - -Logging levels are, in priority order of highest to lowest: - -- `FATAL` - the server is about to exit -- `CRITICAL` - an event occurred that is bad enough to warrant paging someone -- `ERROR` - a user-visible error occurred (though it may be limited in scope) -- `WARNING` - a coder may want to be notified, but the error was probably not - user-visible -- `INFO` - normal informational messages -- `DEBUG` - coder-level debugging information -- `TRACE` - intensive debugging information - -Each node may also optionally choose to *not* pass messages up to the parent -node. - -The `LoggerFactory` builder is used to configure individual log nodes, by -filling in fields and calling the `apply` method. For example, to configure -the root logger to filter at `INFO` level and write to a file: - -```scala -import com.twitter.logging._ - -val factory = LoggerFactory( - node = "", - level = Some(Level.INFO), - handlers = List( - FileHandler( - filename = "/var/log/example/example.log", - rollPolicy = Policy.SigHup - ) - ) -) - -val logger = factory() -``` - -As many `LoggerFactory`s can be configured as you want, so you can attach to -several nodes if you like. To remove all previous configurations, use: - -```scala -Logger.clearHandlers() -``` - -## Handlers - -- `QueueingHandler` - - Queues log records and publishes them in another thread thereby enabling "async logging". - -- `ConsoleHandler` - - Logs to the console. - -- `FileHandler` - - Logs to a file, with an optional file rotation policy. The policies are: - - - `Policy.Never` - always use the same logfile (default) - - `Policy.Hourly` - roll to a new logfile at the top of every hour - - `Policy.Daily` - roll to a new logfile at midnight every night - - `Policy.Weekly(n)` - roll to a new logfile at midnight on day N (0 = Sunday) - - `Policy.SigHup` - reopen the logfile on SIGHUP (for logrotate and similar services) - - When a logfile is rolled, the current logfile is renamed to have the date - (and hour, if rolling hourly) attached, and a new one is started. So, for - example, `test.log` may become `test-20080425.log`, and `test.log` will be - reopened as a new file. - -- `SyslogHandler` - - Log to a syslog server, by host and port. - -- `ScribeHandler` - - Log to a scribe server, by host, port, and category. Buffering and backoff - can also be configured: You can specify how long to collect log lines - before sending them in a single burst, the maximum burst size, and how long - to backoff if the server seems to be offline. - -- `ThrottledHandler` - - Wraps another handler, tracking (and squelching) duplicate messages. If you - use a format string like `"Error %d at %s"`, the log messages will be - de-duped based on the format string, even if they have different parameters. - - -## Formatters - -Handlers usually have a formatter attached to them, and these formatters -generally just add a prefix containing the date, log level, and logger name. - -- `Formatter` - - A standard log prefix like `"ERR [20080315-18:39:05.033] jobs: "`, which - can be configured to truncate log lines to a certain length, limit the lines - of an exception stack trace, and use a special time zone. - - You can override the format string used to generate the prefix, also. - -- `BareFormatterConfig` - - No prefix at all. May be useful for logging info destined for scripts. - -- `SyslogFormatterConfig` - - A formatter required by the syslog protocol, with configurable syslog - priority and date format. - -## Future interrupts - -Method `raise` on `Future` (`def raise(cause: Throwable)`) raises the interrupt described by `cause` to the producer of this `Future`. Interrupt handlers are installed on a `Promise` using `setInterruptHandler`, which takes a partial function: - -```scala -val p = new Promise[T] -p.setInterruptHandler { - case exc: MyException => - // deal with interrupt.. -} -``` - -Interrupts differ in semantics from cancellation in important ways: there can only be one interrupt handler per promise, and interrupts are only delivered if the promise is not yet complete. - -## Time and Duration +# Time and Duration Like arithmetic on doubles, `Time` and `Duration` arithmetic is now free of overflows. Instead, they overflow to `Top` and `Bottom` values, which are analogous to positive and negative infinity. @@ -391,3 +199,9 @@ which is read by applying `elapsed`: ```scala val duration: Duration = elapsed() ``` + +## License + +Copyright 2010 Twitter, Inc. + +Licensed under the Apache License, Version 2.0: https://www.apache.org/licenses/LICENSE-2.0 diff --git a/build.sbt b/build.sbt index 01542f6d5a..ca8c9f4f31 100644 --- a/build.sbt +++ b/build.sbt @@ -1,72 +1,177 @@ -import scoverage.ScoverageKeys +Global / onChangedBuildSource := ReloadOnSourceChanges +Global / excludeLintKeys += scalacOptions // might be actually unused in util-doc module but not sure // All Twitter library releases are date versioned as YY.MM.patch -val releaseVersion = "18.9.0" +val releaseVersion = "24.8.0-SNAPSHOT" -val zkVersion = "3.5.0-alpha" +val slf4jVersion = "1.7.30" +val jacksonVersion = "2.14.3" +val json4sVersion = "4.0.3" +val mockitoVersion = "3.3.3" +val mockitoScalaVersion = "1.14.8" + +val zkVersion = "3.5.6" val zkClientVersion = "0.0.81" val zkGroupVersion = "0.0.92" -val zkDependency = "org.apache.zookeeper" % "zookeeper" % zkVersion excludeAll( +val zkDependency = "org.apache.zookeeper" % "zookeeper" % zkVersion excludeAll ( ExclusionRule("com.sun.jdmk", "jmxtools"), ExclusionRule("com.sun.jmx", "jmxri"), ExclusionRule("javax.jms", "jms") ) -val slf4jVersion = "1.7.21" -val jacksonVersion = "2.9.6" -val guavaLib = "com.google.guava" % "guava" % "19.0" -val caffeineLib = "com.github.ben-manes.caffeine" % "caffeine" % "2.3.4" +val guavaLib = "com.google.guava" % "guava" % "25.1-jre" +val caffeineLib = "com.github.ben-manes.caffeine" % "caffeine" % "2.9.3" val jsr305Lib = "com.google.code.findbugs" % "jsr305" % "2.0.1" -val scalacheckLib = "org.scalacheck" %% "scalacheck" % "1.13.4" % "test" +val scalacheckLib = "org.scalacheck" %% "scalacheck" % "1.15.4" % "test" val slf4jApi = "org.slf4j" % "slf4j-api" % slf4jVersion -val slf4jSimple = "org.slf4j" % "slf4j-simple" % slf4jVersion +val snakeyaml = "org.yaml" % "snakeyaml" % "1.28" -val defaultProjectSettings = Seq( - scalaVersion := "2.12.4", - crossScalaVersions := Seq("2.11.11", "2.12.4") +def travisTestJavaOptions: Seq[String] = { + // We have some custom configuration for the Travis environment + // https://docs.travis-ci.com/user/environment-variables/#default-environment-variables + val travisBuild = sys.env.getOrElse("TRAVIS", "false").toBoolean + if (travisBuild) { + Seq( + "-DSKIP_FLAKY=true", + "-DSKIP_FLAKY_TRAVIS=true" + ) + } else { + Seq( + "-DSKIP_FLAKY=true" + ) + } +} + +def gcJavaOptions: Seq[String] = { + val javaVersion = System.getProperty("java.version") + if (javaVersion.startsWith("1.8")) { + jdk8GcJavaOptions + } else { + jdk11GcJavaOptions + } +} + +def jdk8GcJavaOptions: Seq[String] = { + Seq( + "-XX:+UseParNewGC", + "-XX:+UseConcMarkSweepGC", + "-XX:+CMSParallelRemarkEnabled", + "-XX:+CMSClassUnloadingEnabled", + "-XX:ReservedCodeCacheSize=128m", + "-XX:SurvivorRatio=128", + "-XX:MaxTenuringThreshold=0", + "-Xss8M", + "-Xms512M", + "-Xmx2G" + ) +} + +def jdk11GcJavaOptions: Seq[String] = { + Seq( + "-XX:+UseConcMarkSweepGC", + "-XX:+CMSParallelRemarkEnabled", + "-XX:+CMSClassUnloadingEnabled", + "-XX:ReservedCodeCacheSize=128m", + "-XX:SurvivorRatio=128", + "-XX:MaxTenuringThreshold=0", + "-Xss8M", + "-Xms512M", + "-Xmx2G" + ) +} + +val _scalaVersion = "2.13.6" +val _crossScalaVersions = Seq("2.12.12", "2.13.6") + +val defaultScalaSettings = Seq( + scalaVersion := _scalaVersion, + crossScalaVersions := _crossScalaVersions +) +val defaultScala3EnabledSettings = Seq( + scalaVersion := _scalaVersion, + crossScalaVersions := _crossScalaVersions ++ Seq("3.0.2") +) + +// Our dependencies or compiler options may differ for both Scala 2 and 3. We branch here +// to account for there differences but should merge these artifacts as they are updated. +val scalaDependencies = Seq("org.scala-lang.modules" %% "scala-collection-compat" % "2.4.4") +val scala2cOptions = Seq( + "-target:jvm-1.8", + // Needs -missing-interpolator due to https://issues.scala-lang.org/browse/SI-8761 + "-Xlint:-missing-interpolator", + "-Yrangepos" +) +val scala3cOptions = Seq( + "-Xtarget:8" +) + +val scala3Dependencies = scalaDependencies ++ Seq( + "org.scalatest" %% "scalatest" % "3.2.9" % "test", + "org.scalatestplus" %% "junit-4-13" % "3.2.9.0" % "test" +) +val scala2Dependencies = scalaDependencies ++ Seq( + "org.scalatest" %% "scalatest" % "3.1.2" % "test", + "org.scalatestplus" %% "junit-4-12" % "3.1.2.0" % "test" ) val baseSettings = Seq( version := releaseVersion, organization := "com.twitter", // Workaround for a scaladoc bug which causes it to choke on empty classpaths. - unmanagedClasspath in Compile += Attributed.blank(new java.io.File("doesnotexist")), + Compile / unmanagedClasspath += Attributed.blank(new java.io.File("doesnotexist")), libraryDependencies ++= Seq( - // See http://www.scala-sbt.org/0.13/docs/Testing.html#JUnit + // See https://www.scala-sbt.org/0.13/docs/Testing.html#JUnit "com.novocode" % "junit-interface" % "0.11" % "test", - "org.mockito" % "mockito-all" % "1.10.19" % "test", - "org.scalatest" %% "scalatest" % "3.0.0" % "test" ), - - ScoverageKeys.coverageHighlighting := true, + Test / fork := true, // We have to fork to get the JavaOptions + // Workaround for cross building HealthyQueue.scala, which is not compatible between + // 2.12-, 2.13+, and scala 3. + Compile / unmanagedSourceDirectories += { + val sourceDir = (Compile / sourceDirectory).value + CrossVersion.partialVersion(scalaVersion.value) match { + case Some((2, n)) if n < 13 => sourceDir / "scala-2.12-" + case _ => sourceDir / "scala-2.13+" + } + }, + // Let's skip compiling docs for Scala 3 due a Dotty Scaladoc bug where generating the docs + // requires network access. See https://github.com/lampepfl/dotty/issues/13272 or CSL-11251. + Compile / doc / sources := { + if (scalaVersion.value.startsWith("3")) Seq.empty + else (Compile / doc / sources).value + }, resolvers += "Sonatype OSS Snapshots" at "https://oss.sonatype.org/content/repositories/snapshots", - scalacOptions := Seq( - // Note: Add -deprecation when deprecated methods are removed - "-target:jvm-1.8", "-unchecked", + "-deprecation", "-feature", - "-encoding", "utf8", - // Needs -missing-interpolator due to https://issues.scala-lang.org/browse/SI-8761 - "-Xlint:-missing-interpolator" - ), - + "-encoding", + "utf8", + ) ++ { + if (scalaVersion.value.startsWith("2")) scala2cOptions + else scala3cOptions + }, // Note: Use -Xlint rather than -Xlint:unchecked when TestThriftStructure // warnings are resolved javacOptions ++= Seq("-Xlint:unchecked", "-source", "1.8", "-target", "1.8"), - javacOptions in doc := Seq("-source", "1.8"), - + doc / javacOptions := Seq("-source", "1.8"), + javaOptions ++= Seq( + "-Djava.net.preferIPv4Stack=true", + "-XX:+AggressiveOpts", + "-server" + ), + javaOptions ++= gcJavaOptions, + Test / javaOptions ++= travisTestJavaOptions, // -a: print stack traces for failing asserts testOptions += Tests.Argument(TestFrameworks.JUnit, "-a"), - // This is bad news for things like com.twitter.util.Time - parallelExecution in Test := false, - + Test / parallelExecution := false, // Sonatype publishing - publishArtifact in Test := false, + Test / publishArtifact := false, pomIncludeRepository := { _ => false }, publishMavenStyle := true, + publishConfiguration := publishConfiguration.value.withOverwrite(true), + publishLocalConfiguration := publishLocalConfiguration.value.withOverwrite(true), autoAPIMappings := true, apiURL := Some(url("https://twitter.github.io/util/docs/")), pomExtra := @@ -74,7 +179,7 @@ val baseSettings = Seq( Apache License, Version 2.0 - http://www.apache.org/licenses/LICENSE-2.0 + https://www.apache.org/licenses/LICENSE-2.0 @@ -93,353 +198,547 @@ val baseSettings = Seq( if (version.value.trim.endsWith("SNAPSHOT")) Some("snapshots" at nexus + "content/repositories/snapshots") else - Some("releases" at nexus + "service/local/staging/deploy/maven2") + Some("releases" at nexus + "service/local/staging/deploy/maven2") } ) -val sharedSettings = defaultProjectSettings ++ baseSettings +val sharedSettings = + baseSettings ++ + defaultScalaSettings ++ + Seq(libraryDependencies ++= scala2Dependencies) + +val sharedScala3EnabledSettings = + baseSettings ++ + defaultScala3EnabledSettings ++ + Seq(libraryDependencies ++= scala3Dependencies) ++ + Seq(excludeDependencies ++= Seq("org.scala-lang.modules" % "scala-collection-compat_3")) + +// settings for projects that are cross compiled with scala 2.10 +val settingsCrossCompiledWithTwoTen = + baseSettings ++ + Seq( + crossScalaVersions := Seq("2.10.7") ++ _crossScalaVersions, + scalaVersion := _scalaVersion, + javacOptions ++= Seq("-source", "1.8", "-target", "1.8", "-Xlint:unchecked"), + doc / javacOptions := Seq("-source", "1.8"), + libraryDependencies ++= Seq( + "org.scalacheck" %% "scalacheck" % "1.14.3" % "test" + ) + ) lazy val noPublishSettings = Seq( - publish := {}, - publishLocal := {}, - publishArtifact := false, - // sbt-pgp's publishSigned task needs this defined even though it is not publishing. - publishTo := Some(Resolver.file("Unused transient repository", file("target/unusedrepo"))) + publish / skip := true ) +def scalatestMockitoVersionedDep(scalaVersion: String) = { + if (scalaVersion.startsWith("2")) { + Seq( + "org.scalatestplus" %% "mockito-3-3" % "3.1.2.0" % "test" + ) + } else { + Seq( + "org.scalatestplus" %% "mockito-3-4" % "3.2.9.0" % "test" + ) + } +} + +def scalatestScalacheckVersionedDep(scalaVersion: String) = { + if (scalaVersion.startsWith("2")) { + Seq( + "org.scalatestplus" %% "scalacheck-1-14" % "3.1.2.0" % "test" + ) + } else { + Seq( + "org.scalatestplus" %% "scalacheck-1-15" % "3.2.9.0" % "test" + ) + } +} + lazy val util = Project( id = "util", base = file(".") ).enablePlugins( - ScalaUnidocPlugin -).settings( - sharedSettings ++ - noPublishSettings ++ - Seq( - unidocProjectFilter in (ScalaUnidoc, unidoc) := inAnyProject -- inProjects(utilBenchmark) - ) -).aggregate( - utilApp, - utilBenchmark, - utilCache, - utilCacheGuava, - utilClassPreloader, - utilCodec, - utilCollection, - utilCore, - utilDoc, - utilFunction, - utilHashing, - utilJvm, - utilLint, - utilLogging, - utilReflect, - utilRegistry, - utilSecurity, - utilSlf4jApi, - utilSlf4jJulBridge, - utilStats, - utilTest, - utilThrift, - utilTunable, - utilZk, - utilZkTest -) + ScalaUnidocPlugin + ).settings( + sharedSettings ++ + noPublishSettings ++ + Seq( + ScalaUnidoc / unidoc / unidocProjectFilter := inAnyProject -- inProjects(utilBenchmark) + ) + ).aggregate( + utilApp, + utilAppLifecycle, + utilBenchmark, + utilCache, + utilCacheGuava, + utilClassPreloader, + utilCodec, + utilCore, + utilDoc, + utilFunction, + utilHashing, + utilJacksonAnnotations, + utilJackson, + utilJvm, + utilLint, + utilLogging, + utilMock, + utilReflect, + utilRegistry, + utilRouting, + utilSecurity, + utilSecurityTestCerts, + utilSlf4jApi, + utilSlf4jJulBridge, + utilStats, + utilTest, + utilThrift, + utilTunable, + utilValidator, + utilValidatorConstraint, + utilZk, + utilZkTest + ) lazy val utilApp = Project( id = "util-app", base = file("util-app") ).settings( - sharedSettings + sharedScala3EnabledSettings + ).settings( + name := "util-app", + libraryDependencies ++= Seq( + "org.mockito" % "mockito-core" % mockitoVersion % "test", + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ).dependsOn(utilAppLifecycle, utilCore, utilRegistry) + +lazy val utilAppLifecycle = Project( + id = "util-app-lifecycle", + base = file("util-app-lifecycle") ).settings( - name := "util-app" -).dependsOn(utilCore, utilRegistry) + sharedScala3EnabledSettings + ).settings( + name := "util-app-lifecycle" + ).dependsOn(utilCore) lazy val utilBenchmark = Project( id = "util-benchmark", base = file("util-benchmark") ).settings( - sharedSettings -).enablePlugins( - JmhPlugin -).settings( - name := "util-benchmark" -).dependsOn(utilCore, utilHashing, utilJvm, utilStats) + sharedSettings + ).enablePlugins( + JmhPlugin + ).settings( + name := "util-benchmark" + ).dependsOn( + utilCore, + utilHashing, + utilJackson, + utilJvm, + utilReflect, + utilStats, + utilValidator + ) lazy val utilCache = Project( id = "util-cache", base = file("util-cache") ).settings( - sharedSettings -).settings( - name := "util-cache", - libraryDependencies ++= Seq(caffeineLib, jsr305Lib) -).dependsOn(utilCore) + sharedScala3EnabledSettings + ).settings( + name := "util-cache", + libraryDependencies ++= Seq( + caffeineLib, + jsr305Lib, + "org.mockito" % "mockito-core" % mockitoVersion % "test", + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ).dependsOn(utilCore) lazy val utilCacheGuava = Project( id = "util-cache-guava", base = file("util-cache-guava") ).settings( - sharedSettings -).settings( - name := "util-cache-guava", - libraryDependencies ++= Seq(guavaLib, jsr305Lib) -).dependsOn(utilCache % "compile->compile;test->test", utilCore) + sharedScala3EnabledSettings + ).settings( + name := "util-cache-guava", + libraryDependencies ++= Seq(guavaLib, jsr305Lib) + ).dependsOn(utilCache % "compile->compile;test->test", utilCore) lazy val utilClassPreloader = Project( id = "util-class-preloader", base = file("util-class-preloader") ).settings( - sharedSettings -).settings( - name := "util-class-preloader" -).dependsOn(utilCore) + sharedSettings + ).settings( + name := "util-class-preloader" + ).dependsOn(utilCore) lazy val utilCodec = Project( id = "util-codec", base = file("util-codec") ).settings( - sharedSettings -).settings( - name := "util-codec" -).dependsOn(utilCore) - -lazy val utilCollection = Project( - id = "util-collection", - base = file("util-collection") -).settings( - sharedSettings -).settings( - name := "util-collection", - libraryDependencies ++= Seq( - scalacheckLib - ) -).dependsOn(utilCore % "compile->compile;test->test") + sharedScala3EnabledSettings + ).settings( + name := "util-codec" + ).dependsOn(utilCore) lazy val utilCore = Project( id = "util-core", base = file("util-core") ).settings( - sharedSettings -).settings( - name := "util-core", - libraryDependencies ++= Seq( - caffeineLib % "test", - scalacheckLib, - "org.scala-lang" % "scala-reflect" % scalaVersion.value, - "org.scala-lang.modules" %% "scala-parser-combinators" % "1.0.4" - ), - resourceGenerators in Compile += Def.task { - val projectName = name.value - val file = resourceManaged.value / "com" / "twitter" / projectName / "build.properties" - val buildRev = scala.sys.process.Process("git" :: "rev-parse" :: "HEAD" :: Nil).!!.trim - val buildName = new java.text.SimpleDateFormat("yyyyMMdd-HHmmss").format(new java.util.Date) - val contents = s"name=$projectName\nversion=${version.value}\nbuild_revision=$buildRev\nbuild_name=$buildName" - IO.write(file, contents) - Seq(file) - }.taskValue, - sources in (Compile, doc) := { - // Monitors.java causes "not found: type Monitor$" (CSL-5034) - // so exclude it from the sources used for scaladoc - val previous = (sources in (Compile, doc)).value - previous.filterNot(file => file.getName() == "Monitors.java") - } -).dependsOn(utilFunction) + sharedScala3EnabledSettings + ).settings( + name := "util-core", + // Moved some code to 'concurrent-extra' to conform to Pants' 1:1:1 principle (https://www.pantsbuild.org/build_files.html#target-granularity) + // so that util-core would work better for Pants projects in IntelliJ. + Compile / unmanagedSourceDirectories += baseDirectory.value / "concurrent-extra", + libraryDependencies ++= Seq( + caffeineLib % "test", + scalacheckLib, + "org.mockito" % "mockito-core" % mockitoVersion % "test", + ) ++ + scalatestMockitoVersionedDep(scalaVersion.value) ++ + scalatestScalacheckVersionedDep(scalaVersion.value) ++ { + if (scalaVersion.value.startsWith("2")) { + Seq( + "org.scala-lang" % "scala-reflect" % scalaVersion.value, + "org.scala-lang.modules" %% "scala-parser-combinators" % "1.1.2" + ) + } else { + Seq( + "org.scala-lang.modules" %% "scala-parser-combinators" % "2.0.0" + ) + } + }, + Compile / resourceGenerators += Def.task { + val projectName = name.value + val file = resourceManaged.value / "com" / "twitter" / projectName / "build.properties" + val buildRev = scala.sys.process.Process("git" :: "rev-parse" :: "HEAD" :: Nil).!!.trim + val buildName = new java.text.SimpleDateFormat("yyyyMMdd-HHmmss").format(new java.util.Date) + val contents = + s"name=$projectName\nversion=${version.value}\nbuild_revision=$buildRev\nbuild_name=$buildName" + IO.write(file, contents) + Seq(file) + }.taskValue, + Compile / doc / sources := { + // Monitors.java causes "not found: type Monitor$" (CSL-5034) + // so exclude it from the sources used for scaladoc + val previous = (Compile / doc / sources).value + previous.filterNot(file => file.getName() == "Monitors.java") + } + ).dependsOn(utilFunction) lazy val utilDoc = Project( id = "util-doc", base = file("doc") ).enablePlugins( - SphinxPlugin -).settings( - sharedSettings ++ - noPublishSettings ++ Seq( - scalacOptions in doc ++= Seq("-doc-title", "Util", "-doc-version", version.value), - includeFilter in Sphinx := ("*.html" | "*.png" | "*.svg" | "*.js" | "*.css" | "*.gif" | "*.txt") - )) + SphinxPlugin + ).settings( + sharedSettings ++ + noPublishSettings ++ Seq( + doc / scalacOptions ++= Seq("-doc-title", "Util", "-doc-version", version.value), + Sphinx / includeFilter := ("*.html" | "*.png" | "*.svg" | "*.js" | "*.css" | "*.gif" | "*.txt") + )) lazy val utilFunction = Project( id = "util-function", base = file("util-function") ).settings( - sharedSettings -).settings( - name := "util-function" -) + sharedSettings + ).settings( + name := "util-function" + ) lazy val utilReflect = Project( id = "util-reflect", base = file("util-reflect") ).settings( - sharedSettings -).settings( - name := "util-reflect", - libraryDependencies ++= Seq( - "asm" % "asm" % "3.3.1", - "asm" % "asm-util" % "3.3.1", - "asm" % "asm-commons" % "3.3.1", - "cglib" % "cglib" % "2.2.2" - ) -).dependsOn(utilCore) + sharedSettings + ).settings( + name := "util-reflect" + ).dependsOn(utilCore) lazy val utilHashing = Project( id = "util-hashing", base = file("util-hashing") ).settings( - sharedSettings -).settings( - name := "util-hashing", - libraryDependencies += scalacheckLib -).dependsOn(utilCore % "test") + sharedScala3EnabledSettings + ).settings( + name := "util-hashing", + libraryDependencies ++= Seq( + scalacheckLib + ) ++ scalatestScalacheckVersionedDep(scalaVersion.value) + ).dependsOn(utilCore % "test") + +lazy val utilJacksonAnnotations = Project( + id = "util-jackson-annotations", + base = file("util-jackson-annotations") +).settings( + sharedSettings + ).settings( + name := "util-jackson-annotations" + ) -lazy val utilIntellij = Project( - id = "util-intellij", - base = file("util-intellij") -).settings( - baseSettings -).settings( - name := "util-intellij", - scalaVersion := "2.11.11" -).dependsOn(utilCore % "test") +lazy val utilJackson = Project( + id = "util-jackson", + base = file("util-jackson") +).settings( + sharedSettings + ).settings( + name := "util-jackson", + libraryDependencies ++= Seq( + "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, + "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion, + "com.fasterxml.jackson.dataformat" % "jackson-dataformat-yaml" % jacksonVersion, + "com.fasterxml.jackson.datatype" % "jackson-datatype-jsr310" % jacksonVersion, + "com.fasterxml.jackson.module" %% "jackson-module-scala" % jacksonVersion exclude ("com.google.guava", "guava"), + "jakarta.validation" % "jakarta.validation-api" % "3.0.0", + "org.json4s" %% "json4s-core" % json4sVersion, + "com.fasterxml.jackson.core" % "jackson-annotations" % jacksonVersion % "test", + scalacheckLib, + "org.scalatestplus" %% "scalacheck-1-14" % "3.1.2.0" % "test", + "org.slf4j" % "slf4j-simple" % slf4jVersion % "test" + ) + ).dependsOn( + utilCore, + utilJacksonAnnotations, + utilMock % Test, + utilReflect, + utilSlf4jApi, + utilValidator) lazy val utilJvm = Project( id = "util-jvm", base = file("util-jvm") ).settings( - sharedSettings -).settings( - name := "util-jvm" -).dependsOn(utilApp, utilCore, utilStats, utilTest % "test") + sharedScala3EnabledSettings + ).settings( + name := "util-jvm", + libraryDependencies ++= Seq( + "org.mockito" % "mockito-core" % mockitoVersion % "test", + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ).dependsOn(utilApp, utilCore, utilStats) lazy val utilLint = Project( id = "util-lint", base = file("util-lint") ).settings( - sharedSettings -).settings( - name := "util-lint" -).dependsOn(utilCore) + sharedScala3EnabledSettings + ).settings( + name := "util-lint" + ).dependsOn(utilCore) lazy val utilLogging = Project( id = "util-logging", base = file("util-logging") ).settings( - sharedSettings + sharedScala3EnabledSettings + ).settings( + name := "util-logging" + ).dependsOn(utilCore, utilApp, utilStats) + +lazy val utilMock = Project( + id = "util-mock", + base = file("util-mock") +).settings( + sharedSettings + ).settings( + name := "util-mock", + libraryDependencies ++= Seq( + // only depend on mockito-scala; it will pull in the correct corresponding version of mockito. + "org.mockito" %% "mockito-scala" % mockitoScalaVersion, + "org.mockito" %% "mockito-scala-scalatest" % mockitoScalaVersion % "test" + ) + ).dependsOn(utilCore % "test") + +lazy val utilRegistry = Project( + id = "util-registry", + base = file("util-registry") +).settings( + sharedScala3EnabledSettings + ).settings( + name := "util-registry", + libraryDependencies ++= Seq( + "org.mockito" % "mockito-core" % mockitoVersion % "test", + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ).dependsOn(utilCore) + +lazy val utilRouting = Project( + id = "util-routing", + base = file("util-routing") ).settings( - name := "util-logging" -).dependsOn(utilCore, utilApp, utilStats) + sharedScala3EnabledSettings + ).settings( + name := "util-routing" + ).dependsOn(utilCore, utilSlf4jApi) lazy val utilSlf4jApi = Project( id = "util-slf4j-api", base = file("util-slf4j-api") ).settings( - sharedSettings -).settings( - name := "util-slf4j-api", - libraryDependencies ++= Seq( - slf4jApi, - slf4jSimple - ) -).dependsOn(utilCore % "test") + sharedScala3EnabledSettings + ).settings( + name := "util-slf4j-api", + libraryDependencies ++= Seq( + slf4jApi, + "org.mockito" % "mockito-core" % mockitoVersion % "test", + "org.slf4j" % "slf4j-simple" % slf4jVersion % "test", + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ).dependsOn(utilCore % "test") lazy val utilSlf4jJulBridge = Project( id = "util-slf4j-jul-bridge", base = file("util-slf4j-jul-bridge") ).settings( - sharedSettings -).settings( - name := "util-slf4j-jul-bridge", - libraryDependencies ++= Seq( - slf4jApi, - "org.slf4j" % "jul-to-slf4j" % slf4jVersion) -).dependsOn(utilCore, utilSlf4jApi) - -lazy val utilRegistry = Project( - id = "util-registry", - base = file("util-registry") -).settings( - sharedSettings -).settings( - name := "util-registry" -).dependsOn(utilCore) + sharedScala3EnabledSettings + ).settings( + name := "util-slf4j-jul-bridge", + libraryDependencies ++= Seq(slf4jApi, "org.slf4j" % "jul-to-slf4j" % slf4jVersion) + ).dependsOn(utilCore, utilApp, utilSlf4jApi) lazy val utilSecurity = Project( id = "util-security", base = file("util-security") ).settings( - sharedSettings -).settings( - name := "util-security" -).dependsOn(utilCore, utilLogging) + sharedScala3EnabledSettings + ).settings( + name := "util-security", + libraryDependencies ++= Seq( + scalacheckLib, + snakeyaml + ) ++ scalatestScalacheckVersionedDep(scalaVersion.value) + ).dependsOn(utilCore, utilLogging, utilSecurityTestCerts % "test") + +lazy val utilSecurityTestCerts = Project( + id = "util-security-test-certs", + base = file("util-security-test-certs") +).settings( + sharedSettings + ).settings( + name := "util-security-test-certs" + ) lazy val utilStats = Project( id = "util-stats", base = file("util-stats") ).settings( - sharedSettings -).settings( - name := "util-stats", - libraryDependencies ++= Seq(caffeineLib, jsr305Lib, scalacheckLib) -).dependsOn(utilCore, utilLint) + sharedScala3EnabledSettings + ).settings( + name := "util-stats", + libraryDependencies ++= Seq( + caffeineLib, + jsr305Lib, + scalacheckLib, + "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, + "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion, + ("com.fasterxml.jackson.module" %% "jackson-module-scala" % jacksonVersion exclude ("com.google.guava", "guava")) + .cross(CrossVersion.for3Use2_13), + "org.mockito" % "mockito-core" % mockitoVersion % "test" + ) ++ scalatestMockitoVersionedDep(scalaVersion.value) + ++ scalatestScalacheckVersionedDep(scalaVersion.value) + ++ { + CrossVersion.partialVersion(scalaVersion.value) match { + case Some((2, major)) if major <= 12 => + Seq() + case _ => + Seq("org.scala-lang.modules" %% "scala-parallel-collections" % "1.0.3" % "test") + } + } + ).dependsOn(utilApp, utilCore, utilLint) lazy val utilTest = Project( id = "util-test", base = file("util-test") ).settings( - sharedSettings -).settings( - name := "util-test", - libraryDependencies ++= Seq( - "org.scalatest" %% "scalatest" % "3.0.0", - "org.mockito" % "mockito-all" % "1.10.19" - ) -).dependsOn(utilCore, utilLogging) + sharedSettings + ).settings( + name := "util-test", + libraryDependencies ++= Seq( + "org.mockito" % "mockito-all" % "1.10.19", + "org.scalatestplus" %% "mockito-1-10" % "3.1.0.0" + ) + ).dependsOn(utilCore, utilLogging, utilStats) lazy val utilThrift = Project( id = "util-thrift", base = file("util-thrift") ).settings( - sharedSettings -).settings( - name := "util-thrift", - libraryDependencies ++= Seq( - "org.apache.thrift" % "libthrift" % "0.10.0", - slf4jApi % "provided", - "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, - "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion - ) -).dependsOn(utilCodec) + sharedScala3EnabledSettings + ).settings( + name := "util-thrift", + libraryDependencies ++= Seq( + "org.apache.thrift" % "libthrift" % "0.10.0", + slf4jApi % "provided", + "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, + "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion + ) + ).dependsOn(utilCodec) lazy val utilTunable = Project( id = "util-tunable", base = file("util-tunable") ).settings( - sharedSettings + sharedSettings + ).settings( + name := "util-tunable", + libraryDependencies ++= Seq( + "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, + "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion, + "com.fasterxml.jackson.module" %% "jackson-module-scala" % jacksonVersion exclude ("com.google.guava", "guava") + ) + ).dependsOn(utilApp, utilCore) + +lazy val utilValidator = Project( + id = "util-validator", + base = file("util-validator") +).settings( + sharedSettings + ).settings( + name := "util-validator", + libraryDependencies ++= Seq( + caffeineLib, + scalacheckLib, + "jakarta.validation" % "jakarta.validation-api" % "3.0.0", + "org.hibernate.validator" % "hibernate-validator" % "7.0.1.Final", + "org.glassfish" % "jakarta.el" % "4.0.0", + "org.json4s" %% "json4s-core" % json4sVersion, + "com.fasterxml.jackson.core" % "jackson-annotations" % jacksonVersion % "test", + "org.scalatestplus" %% "scalacheck-1-14" % "3.1.2.0" % "test", + "org.slf4j" % "slf4j-simple" % slf4jVersion % "test" + ) + ).dependsOn(utilCore, utilReflect, utilSlf4jApi, utilValidatorConstraint) + +lazy val utilValidatorConstraint = Project( + id = "util-validator-constraints", + base = file("util-validator-constraints") ).settings( - name := "util-tunable", - libraryDependencies ++= Seq( - "com.fasterxml.jackson.core" % "jackson-core" % jacksonVersion, - "com.fasterxml.jackson.core" % "jackson-databind" % jacksonVersion, - "com.fasterxml.jackson.module" %% "jackson-module-scala" % jacksonVersion exclude("com.google.guava", "guava") + settingsCrossCompiledWithTwoTen + ).settings( + libraryDependencies ++= Seq( + "jakarta.validation" % "jakarta.validation-api" % "3.0.0" + ) ) -).dependsOn(utilApp, utilCore) lazy val utilZk = Project( id = "util-zk", base = file("util-zk") ).settings( - sharedSettings -).settings( - name := "util-zk", - libraryDependencies += zkDependency -).dependsOn(utilCore, utilCollection, utilLogging) + sharedSettings + ).settings( + name := "util-zk", + libraryDependencies ++= Seq( + zkDependency, + "org.mockito" % "mockito-core" % mockitoVersion % "test", + "org.scalatestplus" %% "mockito-3-3" % "3.1.2.0" % "test" + ) + ).dependsOn(utilCore, utilLogging) lazy val utilZkTest = Project( id = "util-zk-test", base = file("util-zk-test") ).settings( - sharedSettings -).settings( - name := "util-zk-test", - libraryDependencies += zkDependency -).dependsOn(utilCore % "test") + sharedScala3EnabledSettings + ).settings( + name := "util-zk-test", + libraryDependencies += zkDependency + ).dependsOn(utilCore % "test") diff --git a/doc/src/sphinx/_templates/sidebarintro.html b/doc/src/sphinx/_templates/sidebarintro.html index aa6e686cbb..586120c874 100644 --- a/doc/src/sphinx/_templates/sidebarintro.html +++ b/doc/src/sphinx/_templates/sidebarintro.html @@ -5,6 +5,6 @@

Useful Links

diff --git a/doc/src/sphinx/conf.py b/doc/src/sphinx/conf.py index 90f5ba19d8..65dc7b6c35 100644 --- a/doc/src/sphinx/conf.py +++ b/doc/src/sphinx/conf.py @@ -29,8 +29,8 @@ 'index_logo': None } -project = u'Util' -copyright = u'{} Twitter, Inc'.format(datetime.datetime.now().year) +project = 'Util' +copyright = '{} Twitter, Inc'.format(datetime.datetime.now().year) htmlhelp_basename = "util" release = sbt_versions.find_release(os.path.abspath('../../../project/Build.scala')) version = sbt_versions.release_to_version(release) @@ -40,13 +40,13 @@ # fall back if theme is not there try: __import__('flask_theme_support') -except ImportError, e: - print '-' * 74 - print 'Warning: Flask themes unavailable. Building with default theme' - print 'If you want the Flask themes, run this command and build again:' - print - print ' git submodule update --init' - print '-' * 74 +except ImportError as e: + print('-' * 74) + print('Warning: Flask themes unavailable. Building with default theme') + print('If you want the Flask themes, run this command and build again:') + print() + print(' git submodule update --init') + print('-' * 74) pygments_style = 'tango' html_theme = 'default' diff --git a/doc/src/sphinx/index.rst b/doc/src/sphinx/index.rst index f91666777c..3380ba6da5 100644 --- a/doc/src/sphinx/index.rst +++ b/doc/src/sphinx/index.rst @@ -13,7 +13,8 @@ Guides util-stats/index util-cookbook/index - util-capturepoints/index + util-jackson/index + util-validator/index Changelog diff --git a/doc/src/sphinx/util-capturepoints/index.rst b/doc/src/sphinx/util-capturepoints/index.rst deleted file mode 100644 index ca2f658195..0000000000 --- a/doc/src/sphinx/util-capturepoints/index.rst +++ /dev/null @@ -1,107 +0,0 @@ -Twitter Futures Capture Points -============================== - -Stack traces produced from asynchronous execution have undesirable -properties, e.g., they are disorganized and often contain frames from -internal libraries that programmers need not know about. This result is -inevitable since not only do different parts of the same program get -executed on different threads, and possibly different cores, but -execution also jumps between different frames. This means that each -thread used to execute a program may produce a unique stack trace. Thus, -an Asynchronous Stack Trace (AST) gives a fragmented window view into -the execution of a program and thus makes it difficult to understand the -flow of code because the causality between frames is lost. - -In response to this phenomenon, IntelliJ 2017.1 introduced a feature -called Capture Points which is built on top of the IntelliJ Debugger. A -capture point is a place in a computer program where the debugger -collects and saves stack frames to be used later when we reach a -specific point in the code and we want to see how we got there. IntelliJ -IDEA does this by substituting part of the call stack with the captured -frame. - -We have written capture points for Twitter Futures in an XML file. Users can -debug their asynchronous code more easily with these capture points. Note that -as of October 2017, the capture points work with only Scala 2.11.11. - -Setup -^^^^^ - -The capture points are defined in -``util/util-intellij/src/main/resources/TwitterFuturesCapturePoints.xml``. -To import them into IntelliJ, - -1. Open IntelliJ. In the menu bar, click IntelliJ IDEA > Preferences. -2. Navigate to Build, Execution, Deployment > Debugger > Async - Stacktraces. -3. Click the Import icon on the bottom bar. Find the XML file and click - OK. - -Use -^^^ - -TL;DR: set a breakpoint where you wish to see the stack trace, debug -your code, and look at the "Frames" tab in the Debugger. Any -asynchronous calls in the stack trace will appear in logical order. If -you wish to clean up the stack trace, click the funnel icon in the top -right, "Hide Frames from Libraries". - -We will illustrate how to use the capture points to assist with -debugging with a small example. The example is located in -``util/util-intellij/src/test/scala/com/twitter/util/capturepoints/Demo.scala``. - -A brief explanation of this test: the test passes a ``Promise[Int]`` -through three methods and then sets the ``Promise``\ ’s value in a -``futurePool``. The calls are asynchronous, but the logical flow of the -test is as follows: - -1. ``test`` block calls ``someBusinessLogic`` -2. ``someBusinessLogic`` calls ``moreBusinessLogic`` -3. ``moreBusinessLogic`` calls ``lordBusinessLogic`` -4. ``lordBusinessLogic`` waits for the ``Promise``\ ’s value to be set -5. The ``Promise``\ ’s value is set in the test block (this could happen - at any time; it is not necessarily step number 5) -6. ``lordBusinessLogic`` returns -7. ``test``\ ’s ``result`` variable is ``4``, and the test passes - -Suppose we wish to inspect the stack trace from inside -``lordBusinessLogic``. If we set a breakpoint at line 47 and then debug -the test with pants in IntelliJ with the capture points disabled, we see -the following stack trace: - -:: - - apply$mcVI$sp:47, Demo$$anonfun$lordBusinessLogic$1 - apply:1820, Future$$anonfun$onSuccess$1 - apply:205, Promise$Monitored - run:532, Promise$$anon$7 - run:198, LocalScheduler$Activation - submit:157, LocalScheduler$Activation - submit:274, LocalScheduler - submit:109, Scheduler$ - runq:522, Promise - updatelfEmpty:880, Promise - update:859, Promise - setValue:835, Promise - apply$mcV$sp:20, Demo$$anonfun$3$$anonfun$apply$1 - apply:15, Try$ - run:140, ExecutorServiceFuturePool$$anon$4 - -The library calls in the stack trace do not help us understand our code. -If we enable the capture points and again debug the test, we see the -following stack trace: - -:: - - apply$mcVI$sp:47, Demo$$anonfun$lordBusinessLogic$1 - apply:1820, Future$$anonfun$onSuccess$1 - onSuccess:1819, Future - lordBusinessLogic:42, Demo - moreBusinessLogic:38, Demo - someBusinessLogic:31, Demo - apply:17, Demo$$anonfun$3 - -This is much cleaner and it includes only calls to functions we wrote. -It lets us clearly see that the test block called ``someBusinessLogic``, -that ``someBusinessLogic`` called ``moreBusinessLogic``, and that -``moreBusinessLogic`` called ``lordBusinessLogic``. diff --git a/doc/src/sphinx/util-cookbook/basics.rst b/doc/src/sphinx/util-cookbook/basics.rst index 94c20e1558..0e1e144416 100644 --- a/doc/src/sphinx/util-cookbook/basics.rst +++ b/doc/src/sphinx/util-cookbook/basics.rst @@ -85,7 +85,7 @@ Scala: .. code-block:: scala - import com.twitter.conversions.time._ + import com.twitter.conversions.DurationOps._ import com.twitter.util.{FuturePool, Time} val time = Time.fromMilliseconds(123456L) @@ -135,7 +135,7 @@ Scala: .. code-block:: scala - import com.twitter.conversions.time._ + import com.twitter.conversions.DurationOps._ import com.twitter.util.{Future, MockTimer, Time} Time.withCurrentTimeFrozen { timeControl => diff --git a/doc/src/sphinx/util-cookbook/futures.rst b/doc/src/sphinx/util-cookbook/futures.rst index a46f43a334..8f1e394acb 100644 --- a/doc/src/sphinx/util-cookbook/futures.rst +++ b/doc/src/sphinx/util-cookbook/futures.rst @@ -11,7 +11,7 @@ Conversions between Twitter's Future and Scala's Future Twitter's ``com.twitter.util.Future`` is similar to, but `predates `_, -`Scala's `_ +`Scala's `_ ``scala.concurrent.Future`` and as such they are not directly compatible (e.g. support for continuation-local variables and interruptibility). @@ -29,7 +29,7 @@ Scala: /** Convert from a Twitter Future to a Scala Future */ implicit class RichTwitterFuture[A](val tf: TwitterFuture[A]) extends AnyVal { - def asScala(implicit e: ExecutionContext): ScalaFuture[A] = { + def asScala: ScalaFuture[A] = { val promise: ScalaPromise[A] = ScalaPromise() tf.respond { case Return(value) => promise.success(value) diff --git a/doc/src/sphinx/util-cookbook/index.rst b/doc/src/sphinx/util-cookbook/index.rst index 11f8334263..45d45fabf6 100644 --- a/doc/src/sphinx/util-cookbook/index.rst +++ b/doc/src/sphinx/util-cookbook/index.rst @@ -14,9 +14,10 @@ This is a cookbook with recipes on how to accomplish many common tasks. basics futures io + reader java Contact ------- -For questions email the `finaglers@twitter.com `_ group. \ No newline at end of file +For questions email the `finaglers@googlegroups.com `_ group. diff --git a/doc/src/sphinx/util-cookbook/reader.rst b/doc/src/sphinx/util-cookbook/reader.rst new file mode 100644 index 0000000000..bf8490598b --- /dev/null +++ b/doc/src/sphinx/util-cookbook/reader.rst @@ -0,0 +1,215 @@ +Reader +====== + +A reader exposes a pull-based API to model a potentially infinite stream of elements. Similar to use an +iterator to iterate through all elements from the stream, users are able to read one element at a time, +all elements are read in the order as they are fed to the reader. + +Interface +--------- + +The most frequently used APIs are: + +.. code-block:: scala + + trait Reader { + def read(): Future[Option[A]] + + def discard(): Unit + + def onClose: Future[StreamTermination] + + final def map[B](f: A => B): Reader[B] + + final def flatMap[B](f: A => Reader[B]): Reader[B] + + def flatten[B](implicit ev: A <:< Reader[B]): Reader[B] + } + +You could find a complete list of APIs from our `Reader Scaladoc `_. +The APIs give you abilities to read from the reader, discard the reader after you are done reading or +you don't want others to read from it. You can use flatMap and map on a reader in a similar way you +use them on `Future` and return a new reader after applying all transforming logic you need. Use +`onClose` to check if the stream is terminated (read until the end of stream, discarded, or encountered +an exception). + +read() +^^^^^^ + +Read the next element of this stream. The returned `Future` will resolve into `Some(e)` when the +element is available, or into `None` when the stream is fully-consumed. Reading from a discarded +reader will always return `ReaderDiscardedException`, reading from a failed reader will always resolve +in the failed Future. + +Multiple outstanding reads are not allowed on the same reader, please avoid calling `read()` again until +the future that the previous read returned has been satisfied. + +Duplicate reads are guaranteed not to happen. When properly using the Reader, it's impossible for a +subsequent read to be satisfied before the first read, since the second read will not be issued until +after the first read has been satisfied. + +discard() +^^^^^^^^^ + +Always discard the reader when its output is no longer required to avoid resource leaks. Reading from a +discarded reader will always resolve in `ReaderDiscardedException`. + +Discarding a fully-consumed reader or a failed reader won't override its stream termination to `Discarded`. +A fully-consumed reader remains at `FullyRead` and a failed reader remains at the failed future that +terminates the stream. + +onClose +^^^^^^^ + +A `Future` that resolves once this reader is fully-consumed (will resolves in `FullyRead`), discarded +(will resolve in `Discarded`), or failed (will resolve in the failed future that terminates the stream). +Once `onClose` is resolved in one stream termination, it won't be overridden to another one. + +This is useful for any extra resource cleanup that you must do once the stream is no longer being used, like +measuring the stream lifetime, counting the number of alive streams, counting the number of failed streams, +calculating stream success rate, etc. + +map +^^^ + +Construct a new Reader by lazily applying `f` to every item read from this Reader, where `f` is the function +that transforms data of type A to B. Note this operation only happens everytime a `read()` is called on the +new Reader. + +All operations of the new Reader will be in sync with self Reader. Discarding one Reader will discard +the other Reader. When one Reader's onClose resolves, the other Reader's onClose will be resolved +immediately with the same value. + +flatMap +^^^^^^^ + +Construct a new Reader (aka flat reader in the rest of the context) by lazily applying `f` to every item +read from self Reader (aka outer reader in the rest of the context), in the order as they are fed to the outer reader, +where `f` is the function that transforms data of type A to B and constructs a Reader[B] (aka inner reader in the rest +of the context). Note this operation only happens when we need to create a new inner reader(explained in the next section). + +When a `read()` operation is called on the flat reader, we will read from the latest constructed inner reader, and +return the value. If the current inner reader is fully-consumed (reading from it returns `None`), instead of +returning `None`, we will apply `f` to create the next inner reader and return the value read from the new inner reader. +If all inner readers are fully-consumed, the flat reader will be fully-consumed. If the current inner reader is failed, +reading from the flat reader will always resolve in the same exception as reading from the failed inner reader, and no new +inner reader will be created. If the current inner reader is discarded, reading from the flat reader will always resolve in +`ReaderDiscardedException` and no new inner readers will be created. + +All operations of the flat reader will be in sync with the outer reader. Discarding one Reader will discard the other Reader. +When one Reader's onClose resolves, the other Reader's onClose will be resolved immediately with the same value. + +However, after creating a flat reader via `flatMap`, the flat reader should be the only channel to interact with the stream, +please avoid manipulating the inner reader or the outer reader by reading from it, discarding it, or failing it. This will +affect the bahivor of the flat reader. + + +flatten +^^^^^^^ + +Converts a `Reader[Reader[B]]` into a `Reader[B]`, this is a convenient abstraction to read from a +stream (Reader, aka outer reader in the rest of the context) of Readers (aka inner reader in the rest of the context) as if +it were a single Reader (aka flat reader in the rest of the context). The behavior to read from the flat reader +is similar to use 2 iterators, with one outer iterator to traverse all inner readers from the stream, another +inner iterator to traverse all elements in the inner reader, and move the outer iterator after the inner iterator +finishes traversing an inner reader (the inner reader is fully-consumed). + +When a `read()` operation is called on the flat reader, we will read from the latest constructed inner reader, and +return the value. If the current inner reader is fully-consumed (reading from it returns `None`), instead of +returning `None`, we will read from the next inner reader and return the value read from the new inner reader. +If all inner readers are fully-consumed, the flat reader will be fully-consumed. If the current inner reader is failed, +reading from the flat reader will always resolve in the same exception as reading from the failed inner reader, and no new +inner reader will be created. If the current inner reader is discarded, reading from the flat reader will always resolve in +`ReaderDiscardedException` and no new inner readers will be created. + +All operations of the flat reader will be in sync with the outer reader. Discarding one Reader will discard +the other Reader. When one Reader's onClose resolves, the other Reader's onClose will be resolved immediately +with the same value. + +However, after creating a flat reader via `flatten`, the flat reader should be the only channel to interact with the stream, +please avoid manipulating the inner readers or the outer reader by reading from it, discarding it, or failing it. This will +affect the bahivor of the flat reader. + +Example Usages +-------------- + +Basic operations: read, discard, onClose +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. code-block:: scala + + import com.twitter.io._ + import com.twitter.util._ + + val r1 = Reader.value(1) // Future(Some(1)) + r1.read() // Future(None) + r1.read() // Future(StreamTermination.FullyRead) + r1.onClose // Future(None) + r1.read() + + val r2 = Reader.value(2) + r2.discard() // Future(StreamTermination.Discarded) + r2.onClose // Future(ReaderDiscardedException) + r2.read() + +Construct a read-loop to consume a reader +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +Given the pull-based API, the consumer is responsible for driving the computation. A very typical +code pattern to consume a reader is to use a read-loop: + +.. code-block:: scala + + def consume[A](r: Reader[A]))(process: A => Future[Unit]): Future[Unit] = { + r.read().flatMap { + case Some(a) => process(a).before(consume(r)(process)) + case None => Future.Done // reached the end of the stream; no need to discard + } + } + +The pattern above leverages `Future` recursion to exert back-pressure by allowing only one +outstanding read. Read must not be issued until the previous read finishes. This would always +ensure a finer grained back-pressure in network systems allowing consumers to artificially slow down +the producers and not rely solely on transport and IO buffering. + +.. NOTE:: + One way to reason about the read-loop idiom is to view it as a subscription to a publisher (reader) + in the Pub/Sub terminology. Perhaps the major difference between readers and traditional publishers, + is readers only allow one subscriber (read-loop). It's generally safer to assume the reader is fully + consumed (stream is exhausted) once its read-loop is run. + +Error handling +^^^^^^^^^^^^^^ + +Given the read-loop above, its returned `Future` could be used to observe both successful and +unsuccessful outcomes. + +.. code-block:: scala + + consume(stream)(processor).respond { + case Return(()) => println("Consumed an entire stream successfully.") + case Throw(e) => println(s"Encountered an error while consuming a stream: $e") + } + +If an exception is thrown during read, the reader's onClose will be resolved with the same exception, so +alternatively, you could handle a stream error like below: + +.. code-block:: scala + + val reader = Reader.value(1) + reader.onClose.respond { + case Return(()) => println("Consumed an entire stream successfully.") + case Throw(e) => println(s"Encountered an error while consuming a stream: $e") + } + +.. NOTE:: + Once failed, a stream can not be restarted such that all future reads will resolve into a failure. + There is no need to discard an already failed stream. + +Resource Safety +^^^^^^^^^^^^^^^ + +One of the important implications of readers, and streams in general, is that they are prone to resource +leaks unless fully consumed,discarded, or failed. Specifically, readers backed with a resource, like a +network connection, **MUST** be discarded unless already consumed (EOF observed), or failed to prevent +connection leaks. diff --git a/doc/src/sphinx/util-jackson/index.rst b/doc/src/sphinx/util-jackson/index.rst new file mode 100644 index 0000000000..9e54d0923a --- /dev/null +++ b/doc/src/sphinx/util-jackson/index.rst @@ -0,0 +1,1090 @@ +.. _util-jackson-index: + +util-jackson Guide +================== + +The library builds upon the excellent |jackson-module-scala|_ for `JSON `_ +support by wrapping the `Jackson ScalaObjectMapper `__ +to provide an API very similar to the `Jerkson `__ `Parser `__. + +Additionally, the library provides a default improved `case class deserializer <#improved-case-class-deserializer>`_ +which *accumulates* deserialization errors while parsing JSON into a `case class` instead of failing-fast, +such that all errors can be reported at once. + +.. admonition :: 🚨 This documentation assumes some level of familiarity with `Jackson JSON Processing `__ + + Specifically Jackson `databinding `__ + and the Jackson `ObjectMapper `__. + + There are several `tutorials `__ (and `external documentation `__) + which may be useful if you are unfamiliar with Jackson. + + Additionally, you may want to familiarize yourself with the `Jackson Annotations `_ + as they allow for finer-grain customization of Jackson databinding. + +Case Classes +------------ + +As mentioned, `Jackson `__ is a JSON processing library. We +generally use Jackson for `databinding `__, +or more specifically: + +- object serialization: converting an object **into a JSON String** and +- object deserialization: converting a JSON String **into an object** + +This library integration is primarily centered around serializing and deserializing Scala +`case classes `__. This is because Scala +`case classes `__ map well to the two JSON +structures [`reference `__]: + +- A collection of name/value pairs. Generally termed an *object*. Here, the name/value pairs are the case field name to field value but can also be an actual Scala `Map[T, U]` as well. +- An ordered list of values. Typically an *array*, *vector*, *list*, or *sequence*, which for case classes can be represented by a Scala `Iterable`. + +Library Features +---------------- + +- Usable as a replacement for the |jackson-module-scala|_ or `Jerkson `__ in many cases. +- A `c.t.util.jackson.ScalaObjectMapper `__ which provides additional Scala friendly methods not found in the |jackson-module-scala|_ `ScalaObjectMapper`. +- A custom `case class deserializer `__ which overcomes some of the limitations in |jackson-module-scala|_. +- Integration with the `util-validator `__ Bean Validation 2.0 (`beanvalidation.org `_) style validations during JSON deserialization. +- Utility for quick JSON parsing with `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__ +- Utility for easily comparing JSON strings. + +Basic Usage +----------- + +Let's assume we have these two case classes: + +.. code-block:: scala + + case class Bar(d: String) + case class Foo(a: String, b: Int, c: Bar) + +To **serialize** a case class into a JSON string, use + +.. code-block:: scala + + ScalaObjectMapper#writeValueAsString(any: Any): String // or + JSON#write(any: Any) + +For example: + +.. code-block:: scala + :emphasize-lines: 10, 15, 30, 33, 38 + + Welcome to Scala 2.12.12 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> val foo = Foo("Hello, World", 42, Bar("Goodbye, World")) + foo: Foo = Foo(Hello, World,42,Bar(Goodbye, World)) + + scala> import com.twitter.util.jackson.JSON + import com.twitter.util.jackson.JSON + + scala> JSON.write(foo) + res3: String = {"a":"Hello, World","b":42,"c":{"d":"Goodbye, World"}} + + scala> // or use the configured "pretty print mapper" + + scala> JSON.prettyPrint(foo) + res4: String = + { + "a" : "Hello, World", + "b" : 42, + "c" : { + "d" : "Goodbye, World" + } + } + + scala> // or use a configured ScalaObjectMapper + + scala> import com.twitter.util.jackson.ScalaObjectMapper + import com.twitter.util.jackson.ScalaObjectMapper + + scala> val mapper = ScalaObjectMapper() + mapper: com.twitter.util.jackson.ScalaObjectMapper = com.twitter.util.jackson.ScalaObjectMapper@490d9c41 + + scala> mapper.writeValueAsString(foo) + res0: String = {"a":"Hello, World","b":42,"c":{"d":"Goodbye, World"}} + + scala> // or use the configured "pretty print mapper" + + scala> mapper.writePrettyString(foo) + res1: String = + { + "a" : "Hello, World", + "b" : 42, + "c" : { + "d" : "Goodbye, World" + } + } + +To **deserialize** a JSON string into a case class, use + +.. code-block:: scala + + ScalaObjectMapper#parse[T](s: String): T // or + JSON#parse[T](s: String): T + +For example, assuming the same `Bar` and `Foo` case classes defined above: + +.. code-block:: scala + :emphasize-lines: 10, 21 + + Welcome to Scala 2.12.12 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> val s = """{"a": "Hello, World", "b": 42, "c": {"d": "Goodbye, World"}}""" + s: String = {"a": "Hello, World", "b": 42, "c": {"d": "Goodbye, World"}} + + scala> import com.twitter.util.jackson.JSON + import com.twitter.util.jackson.JSON + + scala> val foo = JSON.parse[Foo](s) + foo: Option[Foo] = Some(Foo(Hello, World,42,Bar(Goodbye, World))) + + scala> // or use a configured ScalaObjectMapper + + scala> import com.twitter.util.jackson.ScalaObjectMapper + import com.twitter.util.jackson.ScalaObjectMapper + + scala> val mapper = ScalaObjectMapper() + mapper: com.twitter.util.jackson.ScalaObjectMapper = com.twitter.util.jackson.ScalaObjectMapper@3b64f131 + + scala> val foo = mapper.parse[Foo](s) + foo: Foo = Foo(Hello, World,42,Bar(Goodbye, World)) + +.. tip:: + + As seen above you can use the `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__ utility for + general JSON serde operations with a default configured `ScalaObjectMapper <#defaults>`__. + + This can be useful when you do not need to use a specifically configured `ScalaObjectMapper` or + do not wish to perform any Bean Validation 2.0 style validations during JSON deserialization since + the `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__ utility **specifically disables** + validation support on its underlying `c.t.util.jackson.ScalaObjectMapper`. + + See the `documentation <#c-t-util-jackson-json>`__ for more information. + +You can find many examples of using the `ScalaObjectMapper` in the various framework tests: + +- Scala `example `__. +- Java `example `__. + +As mentioned above, there is also a plethora of Jackson `tutorials `__ and `HOW-TOs `__ +available online which provide more in-depth examples of how to use a Jackson `ObjectMapper `__. + +`ScalaObjectMapper` +------------------- + +The |ScalaObjectMapper|_ is a thin wrapper around a configured |jackson-module-scala|_ `ScalaObjectMapper`. +However, the util-jackson |ScalaObjectMapper|_ comes configured with several defaults when instantiated. + +Defaults +~~~~~~~~ + +The following integrations are provided by default when using the |ScalaObjectMapper|_: + +- The Jackson `DefaultScalaModule `__. +- The Jackson `JSR310 `__ `c.f.jackson.datatype.jsr310.JavaTimeModule `__. +- A `LongKeyDeserializer `__ which allows for deserializing Scala Maps with Long keys. +- A `WrappedValueSerializer `__ (more information on "WrappedValues" `here `__). +- Twitter `c.t.util.Time `_ and `c.t.util.Duration `_ serializers [`1 `__, `2 `__] and deserializers [`1 `__, `2 `__]. +- An improved `CaseClassDeserializer `__: see details `below <#improved-case-class-deserializer>`__. +- Integration with the `util-validator `__ Bean Validation 2.0 style validations during JSON deserialization. + +Instantiation +~~~~~~~~~~~~~ + +Instantiation of a new |ScalaObjectMapper|_ can done either via a companion object method or via +the `ScalaObject#Builder` to specify custom configuration. + +`ScalaObjectMapper#apply` +^^^^^^^^^^^^^^^^^^^^^^^^^ + +The companion object defines `apply` methods for creation of a |ScalaObjectMapper|_ configured +with the defaults listed above: + +.. code-block:: scala + + import com.twitter.util.jackson.ScalaObjectMapper + import com.fasterxml.jackson.databind.{ObjectMapper => JacksonObjectMapper} + import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} + + val objectMapper: ScalaObjectMapper = ScalaObjectMapper() + + val underlying: JacksonObjectMapper with JacksonScalaObjectMapper = ??? + val objectMapper: ScalaObjectMapper = ScalaObjectMapper(underlying) + +.. important:: + + The above `#apply` which takes an underlying Jackson `ObjectMapper` will **mutate the configuration** + of the underlying Jackson `ObjectMapper` to apply the default configuration to the given Jackson + `ObjectMapper`. Thus it is not expected that this underlying Jackson `ObjectMapper` be a shared + resource. + +Companion Object Wrappers +^^^^^^^^^^^^^^^^^^^^^^^^^ + +The companion object also defines other methods to easily obtain some specifically configured +|ScalaObjectMapper|_ which wraps an *already configured* Jackson `ObjectMapper`: + +.. code-block:: scala + + import com.twitter.util.jackson.ScalaObjectMapper + import com.fasterxml.jackson.databind.{ObjectMapper => JacksonObjectMapper} + import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} + + val underlying: JacksonObjectMapper with JacksonScalaObjectMapper = ??? + + // different from `#apply(underlying)`. wraps a copy of the given Jackson mapper + // and does not apply any configuration. + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.objectMapper(underlying) + + // merely wraps a copy of the given Jackson mapper that is expected to be configured + // with a YAMLFactory, does not apply any configuration. + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.yamlObjectMapper(underlying) + + // only sets the PropertyNamingStrategy to be PropertyNamingStrategy.LOWER_CAMEL_CASE + // to a copy of the given Jackson mapper, does not apply any other configuration. + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.camelCaseObjectMapper(underlying) + + // only sets the PropertyNamingStrategy to be PropertyNamingStrategy.SNAKE_CASE + // to a copy of the given Jackson mapper, does not apply any other configuration. + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.snakeCaseObjectMapper(underlying) + +.. admonition:: These methods clone the underlying mapper, they do not mutate the given Jaskson mapper. + + Note that these methods will *copy* the underlying Jackson mapper (not mutate it) to apply any + necessary configuration to produce a new |ScalaObjectMapper|_, which in this case is only to + change the `PropertyNamingStrategy` accordingly. No other configuration changes are made + to the copy of the underlying Jackson mapper. + + Specifically, note that the `ScalaObjectMapper.objectMapper(underlying)` **wraps a copy but does not mutate the original** + configuration of the given underlying Jackson `ObjectMapper`. This is different from the + `ScalaObjectMapper(underlying)` which **mutates** the given underlying Jackson `ObjectMapper` to + apply all of the default configuration to produce a `ScalaObjectMapper`. + +`ScalaObjectMapper#Builder` +^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +You can use the `ScalaObjectMapper#Builder` for more advanced control over configuration options for +producing a configured `ScalaObjectMapper`. For example, to create an instance of a |ScalaObjectMapper|_ +with case class validation via the util-validator `ScalaValidator` **disabled**: + +.. code :: scala + + import com.twitter.util.jackson.ScalaObjectMapper + + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.builder.withNoValidation.objectMapper + +See the `Advanced Configuration <#advanced-configuration>`__ section for more information. + +Advanced Configuration +~~~~~~~~~~~~~~~~~~~~~~ + +To apply more custom configuration to create a |ScalaObjectMapper|_, there is a builder for +constructing a customized mapper. + +E.g., to set a `PropertyNamingStrategy` different than the default: + +.. code-block:: scala + :emphasize-lines: 3 + + val objectMapper: ScalaObjectMapper = + ScalaObjectMapper.builder + .withPropertyNamingStrategy(PropertyNamingStrategy.KebabCaseStrategy) + .objectMapper + +Or to set additional modules or configuration: + +.. code-block:: scala + :emphasize-lines: 4, 5, 6, 7 + + val objectMapper: ScalaObjectMapper = + ScalaObjectMapper.builder + .withPropertyNamingStrategy(PropertyNamingStrategy.KebabCaseStrategy) + .withAdditionalJacksonModules(Seq(MySimpleJacksonModule)) + .withAdditionalMapperConfigurationFn( + _.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, true) + ) + .objectMapper + +You can also get a `camelCase`, `snake_case`, or even a YAML configured mapper. + +.. code-block:: scala + :emphasize-lines: 7, 15, 23 + + val camelCaseObjectMapper: ScalaObjectMapper = + ScalaObjectMapper.builder + .withAdditionalJacksonModules(Seq(MySimpleJacksonModule)) + .withAdditionalMapperConfigurationFn( + _.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, true) + ) + .camelCaseObjectMapper + + val snakeCaseObjectMapper: ScalaObjectMapper = + ScalaObjectMapper.builder + .withAdditionalJacksonModules(Seq(MySimpleJacksonModule)) + .withAdditionalMapperConfigurationFn( + _.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, true) + ) + .snakeCaseObjectMapper + + val yamlObjectMapper: ScalaObjectMapper = + ScalaObjectMapper.builder + .withAdditionalJacksonModules(Seq(MySimpleJacksonModule)) + .withAdditionalMapperConfigurationFn( + _.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, true) + ) + .yamlObjectMapper + +Access to the underlying Jackson Object Mapper +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The |ScalaObjectMapper|_ is a thin wrapper around a configured Jackson |jackson-module-scala|_ +`com.fasterxml.jackson.module.scala.ScalaObjectMapper`, thus you can always access the underlying +Jackson object mapper by calling `underlying`: + +.. code-block:: scala + :emphasize-lines: 7 + + import com.fasterxml.jackson.databind.ObjectMapper + import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} + import com.twitter.util.jackson.ScalaObjectMapper + + val objectMapper: ScalaObjectMapper = ??? + + val jacksonObjectMapper: ObjectMapper with JacksonScalaObjectMapper = objectMapper.underlying + +Adding a Custom Serializer or Deserializer +------------------------------------------ + +For more information see the Jackson documentation for +`Custom Serializers `__. + +Add a Jackson Module to a |ScalaObjectMapper|_ +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Follow the steps to create a Jackson Module for the custom serializer or deserializer then register +the module to the underlying Jackson mapper from the |ScalaObjectMapper|_ instance: + +.. code-block:: scala + :emphasize-lines: 13, 14, 15, 43, 44, 45, 46, 47, 48, 52, 53 + + import com.fasterxml.jackson.databind.JsonDeserializer + import com.fasterxml.jackson.databind.deser.Deserializers + import com.fasterxml.jackson.databind.module.SimpleModule + import com.fasterxml.jackson.module.scala.JacksonModule + import com.twitter.util.jackson.ScalaObjectMapper + + // custom deserializer + class FooDeserializer extends JsonDeserializer[Foo] { + override def deserialize(...) + } + + // Jackson SimpleModule for custom deserializer + class FooDeserializerModule extends SimpleModule { + addDeserializer(FooDeserializer) + } + + // custom parameterized deserializer + class MapIntIntDeserializer extends JsonDeserializer[Map[Int, Int]] { + override def deserialize(...) + } + + // custom parameterized deserializer resolver + class MapIntIntDeserializerResolver extends Deserializers.Base { + override def findBeanDeserializer( + javaType: JavaType, + config: DeserializationConfig, + beanDesc: BeanDescription + ): MapIntIntDeserializer = { + if (javaType.isMapLikeType && javaType.hasGenericTypes && hasIntTypes(javaType)) { + new MapIntIntDeserializer + } else null + } + + private[this] def hasIntTypes(javaType: JavaType): Boolean = { + val k = javaType.containedType(0) + val v = javaType.containedType(1) + k.isPrimitive && k.getRawClass == classOf[Integer] && + v.isPrimitive && v.getRawClass == classOf[Integer] + } + } + + // Jackson SimpleModule for custom deserializer + class MapIntIntDeserializerModule extends JacksonModule { + override def getModuleName: String = this.getClass.getName + this += { + _.addDeserializers(new MapIntIntDeserializerResolver) + } + } + + ... + + val mapper: ScalaObjectMapper = ??? + mapper.registerModules(new FooDeserializerModule, new MapIntIntDeserializerModule) + +Improved `case class` deserializer +---------------------------------- + +The library provides a `case class deserializer `__ +which overcomes some limitations in |jackson-module-scala|_: + +- Throws a `JsonMappingException` when required fields are missing from the parsed JSON. +- Uses specified `case class` default values when fields are missing in the incoming JSON. +- Properly deserializes a `Seq[Long]` (see: https://github.com/FasterXML/jackson-module-scala/issues/62). +- Supports `"wrapped values" `__ using `c.t.util.jackson.WrappedValue `_. +- Support for field and method level validations via integration with the `util-validator `__ Bean Validation 2.0 style validations during JSON deserialization. +- Accumulates all JSON deserialization errors (instead of failing fast) in a returned sub-class of `JsonMappingException` (see: `CaseClassMappingException `_). + +The `case class deserializer `__ +is added by default when constructing a new |ScalaObjectMapper|_. + +.. tip:: + + Note: with the |CaseClassDeserializer|_, non-option fields without default values are + **considered required**. If a required field is missing, a `CaseClassMappingException` is thrown. + + JSON `null` values are **not allowed and will be treated as a "missing" value**. If necessary, users + can specify a custom deserializer for a field if they want to be able to parse a JSON `null` into + a Scala `null` type for a field. Define your deserializer, `NullAllowedDeserializer` then annotate + the field with `@JsonDeserialize(using = classOf[NullAllowedDeserializer])`. + +`@JsonCreator` Support +~~~~~~~~~~~~~~~~~~~~~~ + +The |CaseClassDeserializer|_ supports specification of a constructor or static factory +method annotated with the Jackson Annotation, `@JsonCreator `_ +(an annotation for indicating a specific constructor or static factory method to use for +instantiation of the case class during deserialization). + +For example, you can annotate a method on the companion object for the case class as a static +factory for instantiation. Any static factory method to use for instantiation **MUST** be specified +on the companion object for case class: + +.. code-block:: scala + :emphasize-lines: 4 + + case class MySimpleCaseClass(int: Int) + + object MySimpleCaseClass { + @JsonCreator + def apply(s: String): MySimpleCaseClass = MySimpleCaseClass(s.toInt) + } + +Or to specify a secondary constructor to use for case class instantiation: + +.. code-block:: scala + :emphasize-lines: 2 + + case class MyCaseClassWithMultipleConstructors(number1: Long, number2: Long, number3: Long) { + @JsonCreator + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } + } + +.. note:: + + If you define multiple constructors on a case class, it is **required** to annotate one of the + constructors with `@JsonCreator`. + + To annotate the primary constructor (as the syntax can seem non-intuitive because the `()` is + required): + + .. code-block:: scala + :emphasize-lines: 1 + + case class MyCaseClassWithMultipleConstructors @JsonCreator()(number1: Long, number2: Long, number3: Long) { + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } + } + + The parens are needed because the Scala class constructor syntax requires constructor + annotations to have exactly one parameter list, possibly empty. + + If you define multiple case class constructors with no visible `@JsonCreator` constructor or + static factory method via a companion, deserialization will error. + +`@JsonFormat` Support +~~~~~~~~~~~~~~~~~~~~~ + +The |CaseClassDeserializer|_ supports `@JsonFormat`-annotated case class fields to properly +contextualize deserialization based on the values in the annotation. + +A common use case is to be able to support deserializing a JSON string into a "time" representation +class based on a specific pattern independent of the time format configured on the `ObjectMapper` or +even the default format for a given deserializer for the type. + +For instance, the library provides a `specific deserializer `_ +for the `com.twitter.util.Time `_ +class. This deserializer is a Jackson `ContextualDeserializer `_ +and will properly take into account a `@JsonFormat`-annotated field. + +However, the |CaseClassDeserializer|_ is invoked first and acts as a proxy for deserializing the time +value. The case class deserializer properly contextualizes the field for correct deserialization by +the `TimeStringDeserializer`. + +Thus if you had a case class defined: + +.. code-block:: scala + :emphasize-lines: 7 + + import com.fasterxml.jackson.annotation.JsonFormat + import com.twitter.util.Time + + case class Event( + id: Long, + description: String, + @JsonFormat(pattern = "yyyy-MM-dd'T'HH:mm:ss.SSSXXX") when: Time + ) + +The following JSON: + +.. code-block:: json + + { + "id": 42, + "description": "Something happened.", + "when": "2018-09-14T23:20:08.000-07:00" + } + +Will always deserialize properly into the case class regardless of the pattern configured on the +`ObjectMapper` or as the default of a contextualized deserializer: + +.. code-block:: scala + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import com.fasterxml.jackson.annotation.JsonFormat + import com.fasterxml.jackson.annotation.JsonFormat + + scala> import com.twitter.util.Time + import com.twitter.util.Time + + scala> case class Event( + | id: Long, + | description: String, + | @JsonFormat(pattern = "yyyy-MM-dd'T'HH:mm:ss.SSSXXX") when: Time + | ) + defined class Event + + scala> val json = """ + | { + | "id": 42, + | "description": "Something happened.", + | "when": "2018-09-14T23:20:08.000-07:00" + | }""".stripMargin + json: String = + " + { + "id": 42, + "description": "Something happened.", + "when": "2018-09-14T23:20:08.000-07:00" + }" + + scala> import com.twitter.util.jackson.ScalaObjectMapper + import com.twitter.util.jackson.ScalaObjectMapper + + scala> val mapper = ScalaObjectMapper() + mapper: com.twitter.util.jackson.ScalaObjectMapper = com.twitter.util.jackson.ScalaObjectMapper@52dc71b2 + + scala> val event: Event = mapper.parse[Event](json) + event: Event = Event(42,Something happened.,2018-09-15 06:20:08 +0000) + +Jackson InjectableValues Support +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +By default, the library does not configure any `com.fasterxml.jackson.databind.InjectableValues `__ +implementation. + +`@InjectableValue` +^^^^^^^^^^^^^^^^^^ + +It does however provide the `c.t.util.jackson.annotation.InjectableValue `__ annotation which can be used to mark *other* `java.lang.annotation.Annotation` interfaces as annotations +which support case class field injection via Jackson `com.fasterxml.jackson.databind.InjectableValues`. + +That is, users can create custom annotations and annotate them with `@InjectableValue`. This +then allows for a configured Jackson `com.fasterxml.jackson.databind.InjectableValues` implementation +to be able to treat these annotations similar to the `@JacksonInject `__. + +This means a custom Jackson `InjectableValues` implementation can use the `@InjectableValue` marker +annotation to resolve fields annotated with annotations that have the `@InjectableValue` marker +annotation as injectable fields. + +For more information on the Jackson `@JacksonInject` or `c.f.databind.InjectableValues` support see the +tutorial `here `__. + +`Mix-in Annotations `_ +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The Jackson `Mix-in Annotations `_ +provide a way to associate annotations to classes without needing to modify the target classes +themselves. It is intended to help support 3rd party datatypes where the user cannot modify the +sources to add annotations. + +The |CaseClassDeserializer|_ supports Jackson `Mix-in Annotations `_ +for specifying field annotations during deserialization with the `case class deserializer `_. + +For example, to deserialize JSON into the following classes that are not yours to annotate: + +.. code-block:: scala + + case class Point(x: Int, y: Int) { + def area: Int = x * y + } + + case class Points(points: Seq[Point]) + +However, you want to enforce field constraints with `validations <../util-validator/index.html>`_ +during deserialization. You can define a `Mix-in`, + +.. code-block:: scala + + import com.fasterxml.jackson.annotation.JsonIgnore + import jakarta.validation.constraints.{Max, Min} + + trait PointMixIn { + @Min(0) @Max(100) def x: Int + @Min(0) @Max(100) def y: Int + @JsonIgnore def area: Int + } + +Then register this `Mix-in` for the `Point` class type. There are several ways to do this: + +Follow the steps to create a Jackson Module for the `Mix-in` then register the module to the +underlying Jackson mapper from the |ScalaObjectMapper|_ instance: + +.. code-block:: scala + + import com.fasterxml.jackson.databind.module.SimpleModule + import com.twitter.util.jackson.ScalaObjectMapper + + object PointMixInModule extends SimpleModule { + setMixInAnnotation(classOf[Point], classOf[PointMixIn]); + } + + ... + + val objectMapper: ScalaObjectMapper = ??? + objectMapper.registerModule(PointMixInModule) + +Or register the `Mix-in` for the class type directly on the underlying Jackson mapper (without a Jackson Module): + +.. code-block:: scala + + import com.twitter.util.jackson.ScalaObjectMapper + + val objectMapper: ScalaObjectMapper = ??? + objectMapper.underlying.addMixin[Point, PointMixIn] + +Deserializing this JSON would then error with failed validations: + +.. code-block:: json + + { + "points": [ + {"x": -1, "y": 120}, + {"x": 4, "y": 99} + ] + } + +.. code-block:: scala + :emphasize-lines: 51, 52, 53 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> case class Point(x: Int, y: Int) { + | def area: Int = x * y + | } + defined class Point + + scala> case class Points(points: Seq[Point]) + defined class Points + + scala> import com.fasterxml.jackson.annotation.JsonIgnore + import com.fasterxml.jackson.annotation.JsonIgnore + + scala> import jakarta.validation.constraints.{Max, Min} + import jakarta.validation.constraints.{Max, Min} + + scala> trait PointMixIn { + | @Min(0) @Max(100) def x: Int + | @Min(0) @Max(100) def y: Int + | @JsonIgnore def area: Int + | } + defined trait PointMixIn + + scala> import com.twitter.util.jackson.ScalaObjectMapper + import com.twitter.util.jackson.ScalaObjectMapper + + scala> val objectMapper: ScalaObjectMapper = ScalaObjectMapper() + objectMapper: com.twitter.util.jackson.ScalaObjectMapper = com.twitter.util.jackson.ScalaObjectMapper@2389f546 + + scala> objectMapper.underlying.addMixin[Point, PointMixIn] + res0: com.fasterxml.jackson.databind.ObjectMapper = com.twitter.util.jackson.ScalaObjectMapper$Builder$$anon$1@5ae22651 + + scala> val json = """ + | { + | "points": [ + | {"x": -1, "y": 120}, + | {"x": 4, "y": 99} + | ] + | }""".stripMargin + json: String = + " + { + "points": [ + {"x": -1, "y": 120}, + {"x": 4, "y": 99} + ] + }" + + scala> val points = objectMapper.parse[Points](json) + com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException: 2 errors encountered during deserialization. + Errors: com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException: points.x: must be greater than or equal to 0 + com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException: points.y: must be less than or equal to 100 + at com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException$.apply(CaseClassMappingException.scala:21) + at com.twitter.util.jackson.caseclass.CaseClassDeserializer.deserialize(CaseClassDeserializer.scala:431) + at com.twitter.util.jackson.caseclass.CaseClassDeserializer.deserializeNonWrapperClass(CaseClassDeserializer.scala:408) + at com.twitter.util.jackson.caseclass.CaseClassDeserializer.deserialize(CaseClassDeserializer.scala:373) + at com.fasterxml.jackson.databind.ObjectMapper._readMapAndClose(ObjectMapper.java:4524) + at com.fasterxml.jackson.databind.ObjectMapper.readValue(ObjectMapper.java:3466) + at com.fasterxml.jackson.module.scala.ScalaObjectMapper.readValue(ScalaObjectMapper.scala:191) + at com.fasterxml.jackson.module.scala.ScalaObjectMapper.readValue$(ScalaObjectMapper.scala:190) + at com.twitter.util.jackson.ScalaObjectMapper$Builder$$anon$1.readValue(ScalaObjectMapper.scala:382) + at com.twitter.util.jackson.ScalaObjectMapper.parse(ScalaObjectMapper.scala:463) + ... 34 elided + +As the first `Point` instance has an x-value less than the minimum of 0 and a y-value greater than +the maximum of 100. + +Known `CaseClassDeserializer` Limitations +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The |CaseClassDeserializer|_ provides a fair amount of utility but can not and does not +support all Jackson Annotations. The behavior of supporting a Jackson Annotation can at times be +ambiguous (or even nonsensical), especially when it comes to combining Jackson Annotations and +injectable field annotations. + +Java Enums +~~~~~~~~~~ + +We recommend the use of `Java Enums `__ +for representing enumerations since they integrate well with Jackson's ObjectMapper and have +exhaustiveness checking as of Scala 2.10. + +The following `Jackson annotations `__ may be +useful when working with Enums: + +- `@JsonValue`: can be used for an overridden `toString` method. +- `@JsonEnumDefaultValue`: can be used for defining a default value when deserializing unknown Enum values. Note that this requires `READ_UNKNOWN_ENUM_VALUES_USING_DEFAULT_VALUE `_ feature to be enabled. + +`c.t.util.jackson.JSON` +----------------------- + +The library provides a utility for default mapping of JSON to an object or writing an object as a JSON +string. This is largely inspired by the `scala.util.parsing.json.JSON `__ +from the Scala Parser Combinators library. + +However, the `c.t.util.jackson.JSON `__ +utility uses a default configured `ScalaObjectMapper <#defaults>`__ and is thus more full featured +than the `scala.util.parsing.json.JSON` utility. + +.. important:: + + The `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__ API does not return an exception when + parsing but rather returns an `Option[T]` result. When parsing is successful, this is a `Some(T)`, + otherwise it is a `None`. But note that the specifics of any failure are lost. + + It is thus also important to note that for this reason that the `c.t.util.jackson.JSON` uses a + `default configured ScalaObjectMapper <#defaults>`__ **with validation specifically disabled**, + such that no `Bean Validation 2.0 <../util-validator>`__ style validations are performed when + parsing with `c.t.util.jackson.JSON`. + + Users should prefer using a configured |ScalaObjectMapper|_ to perform validations in order to + be able to properly handle validation exceptions. + +.. code-block:: scala + :emphasize-lines: 7, 13, 22, 25 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import com.twitter.util.jackson.JSON + import com.twitter.util.jackson.JSON + + scala> val result = JSON.parse[Map[String, Int]]("""{"a": 1, "b": 2}""") + result: Option[Map[String,Int]] = Some(Map(a -> 1, b -> 2)) + + scala> case class FooClass(id: String) + defined class FooClass + + scala> val result = JSON.parse[FooClass]("""{"id": "abcd1234"}""") + result: Option[FooClass] = Some(FooClass(abcd1234)) + + scala> result.get + res0: FooClass = FooClass(abcd1234) + + scala> val f = FooClass("99999999") + f: FooClass = FooClass(99999999) + + scala> JSON.write(f) + res1: String = {"id":"99999999"} + + scala> JSON.prettyPrint(f) + res2: String = + { + "id" : "99999999" + } + + scala> + + +`c.t.util.jackson.YAML` +----------------------- + +Similarly, there is also a utility for `YAML `__ serde operations with many of +the same methods as `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__ using a default |ScalaObjectMapper|_ +configured with a `YAMLFactory`. + +.. code-block:: scala + :emphasize-lines: 8, 9, 10, 11, 12, 18, 19, 20, 29 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import com.twitter.util.jackson.YAML + import com.twitter.util.jackson.YAML + + scala> val result = + | YAML.parse[Map[String, Int]](""" + | |--- + | |a: 1 + | |b: 2 + | |c: 3""".stripMargin) + result: Option[Map[String,Int]] = Some(Map(a -> 1, b -> 2, c -> 3)) + + scala> case class FooClass(id: String) + defined class FooClass + + scala> val result = YAML.parse[FooClass](""" + | |--- + | |id: abcde1234""".stripMargin) + result: Option[FooClass] = Some(FooClass(abcde1234)) + + scala> result.get + res0: FooClass = FooClass(abcde1234) + + scala> val f = FooClass("99999999") + f: FooClass = FooClass(99999999) + + scala> YAML.write(f) + res1: String = + "--- + id: "99999999" + " + + scala> + +.. important:: + + Like with `c.t.util.jackson.JSON <#c-t-util-jackson-json>`__, the `c.t.util.jackson.YAML <#c-t-util-jackson-yaml>`__ + API does not return an exception when parsing but rather returns an `Option[T]` result. + + Users should prefer using a configured |ScalaObjectMapper|_ to perform validations in order to be + able to properly handle validation exceptions. + +`c.t.util.jackson.JsonDiff` +--------------------------- + +The library provides a utility for comparing JSON strings, or structures that can be serialized as +JSON strings, via the `c.t.util.jackson.JsonDiff `__ utility. + +`JsonDiff` provides two functions: `diff` and `assertDiff`. The `diff` method allows the user to +decide how to handle JSON differences by returning an `Option[JsonDiff.Result]` while `assertDiff` +throws an `AssertionError` when a difference is encountered. + +The `JsonDiff.Result#toString` contains a textual representation meant to indicate where the +`expected` and `actual` differ semantically. For this representation, both `expected` and `actual` +are transformed to eliminate insigificant lexical diffences such whitespace, object key ordering, +and escape sequences. Only the first difference in this representation is indicated. +When an `AssertError` is thrown, the `JsonDiff.Result#toString` is used to +populate the exception message. + +For example: + +.. code-block:: scala + :emphasize-lines: 13, 19, 54 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import com.twitter.util.jackson.JsonDiff + import com.twitter.util.jackson.JsonDiff + + scala> val a ="""{"a":1,"b":2}""" + a: String = {"a":1,"b":2} + + scala> val b ="""{"b": 2,"a": 1}""" + b: String = {"b": 2,"a": 1} + + scala> val result = JsonDiff.diff(a, b) // no difference + result: Option[com.twitter.util.jackson.JsonDiff.Result] = None + + scala> val b ="""{"b": 3,"a": 1}""" + b: String = {"b": 3,"a": 1} + + scala> val result = JsonDiff.diff(a, b) // b-values are different + result: Option[com.twitter.util.jackson.JsonDiff.Result] = + Some( * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3}) + + scala> result.get.toString + res1: String = + " * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3}" + + scala> result.get.expected + res3: com.fasterxml.jackson.databind.JsonNode = {"a":1,"b":2} + + scala> result.get.actual + res4: com.fasterxml.jackson.databind.JsonNode = {"b":3,"a":1} + + scala> result.get.expectedPrettyString + res5: String = + { + "a" : 1, + "b" : 2 + } + + scala> result.get.actualPrettyString + res6: String = + { + "b" : 3, + "a" : 1 + } + + scala> import com.twitter.util.Try + import com.twitter.util.Try + + scala> val t = Try(JsonDiff.assertDiff(a, b)) // throws an AssertionError + JSON DIFF FAILED! + * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3} + + scala> t.isThrow + res0: Boolean = true + + scala> t.throwable.getMessage + res1: String = + com.twitter.util.jackson.JsonDiff$ failure + * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3} + + scala> val expected = """{"t1": "24\u00B0C"}""" + expected: String = {"t1": "24\u00B0C"} + + scala> val actual = """{"t1": "24°F", "t2": null}""" + actual: String = {"t1": "24°F", "t2": null} + + scala> val t = Try(JsonDiff.assertDiff(expected, actual)) // throws an AssertionError + JSON DIFF FAILED! + * + Expected: {"t1":"24°C"} + Actual: {"t1":"24°F","t2":null} + + scala> t.isThrow + res0: Boolean = true + + scala> t.throwable.getMessage + res1: String = + com.twitter.util.jackson.JsonDiff$ failure + * + Expected: {"t1":"24°C"} + Actual: {"t1":"24°F","t2":null} + + +Normalization +~~~~~~~~~~~~~ + +Both API methods accept a "normalize function" which is a function to apply on the *actual* to "normalize" any +fields -- such as a timestamp -- before comparing to the *expected*. + +.. code-block:: scala + :emphasize-lines: 43, 44, 45, 46, 49, 52 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import com.twitter.util.jackson.JsonDiff + import com.twitter.util.jackson.JsonDiff + + scala> val a ="""{"a":1,"b":2}""" + a: String = {"a":1,"b":2} + + scala> val b ="""{"b": 3,"a": 1}""" + b: String = {"b": 3,"a": 1} + + scala> val result = JsonDiff.diff(expected = a, actual = b) // b-values are different + result: Option[com.twitter.util.jackson.JsonDiff.Result] = + Some( * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3}) + + scala> result.get.toString + res1: String = + " * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3}" + + scala> JsonDiff.assertDiff(expected = a, actual = b) // throws an AssertionError + JSON DIFF FAILED! + * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3} + java.lang.AssertionError: com.twitter.util.jackson.JsonDiff$ failure + * + Expected: {"a":1,"b":2} + Actual: {"a":1,"b":3} + at com.twitter.util.jackson.JsonDiff$.assert(JsonDiff.scala:206) + at com.twitter.util.jackson.JsonDiff$.assertDiff(JsonDiff.scala:114) + ... 34 elided + + scala> import com.fasterxml.jackson.databind.JsonNode + import com.fasterxml.jackson.databind.JsonNode + + scala> import com.fasterxml.jackson.databind.node.ObjectNode + import com.fasterxml.jackson.databind.node.ObjectNode + + scala> val normalizeFn: JsonNode => JsonNode = { jsonNode: JsonNode => + | jsonNode.asInstanceOf[ObjectNode].put("b", 2) + | } + normalizeFn: com.fasterxml.jackson.databind.JsonNode => com.fasterxml.jackson.databind.JsonNode = $Lambda$1162/1826212603@3da20c42 + + scala> val result = JsonDiff.diff(expected = a, actual = b, normalizeFn) // normalize fn updates b-value of 'actual' to match the 'expected' + result: Option[com.twitter.util.jackson.JsonDiff.Result] = None + + scala> JsonDiff.assertDiff(expected = a, actual = b, normalizeFn) // no exception when 'actual' is normalized + + scala> + +.. |ScalaObjectMapper| replace:: `ScalaObjectMapper` +.. _ScalaObjectMapper: https://github.com/twitter/util/blob/develop/util-jackson/src/main/scala/com/twitter/util/jackson/ScalaObjectMapper.scala + +.. |CaseClassDeserializer| replace:: `util-jackson case class deserializer` +.. _CaseClassDeserializer: https://github.com/twitter/util/blob/develop/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializer.scala + +.. |jackson-module-scala| replace:: `jackson-module-scala` +.. _jackson-module-scala: https://github.com/FasterXML/jackson-module-scala + + diff --git a/doc/src/sphinx/util-stats/faq.rst b/doc/src/sphinx/util-stats/faq.rst index d44b57ba19..874892f641 100644 --- a/doc/src/sphinx/util-stats/faq.rst +++ b/doc/src/sphinx/util-stats/faq.rst @@ -10,15 +10,11 @@ The code was moved to `util `_ in order to allow usage without depending on finagle-core. To ease the transition, the package namespace was left as is. -Why were there three(!) different implementations? +Why were there two(!) different implementations? -------------------------------------------------- -This is somewhat historical, as early on in Twitter, Commons Stats -was created for Java developers and Ostrich was created for Scala developers. -Code initially was written directly to the implementations and there were -(and still are) some subtle differences in their behavior. -``StatsReceiver`` was written in order to bridge the gap between the -two implementations. `finagle-stats` was written quite a bit later in order to -improve upon both of these implementations in terms of +This is somewhat historical, as early on in Twitter, Ostrich was created for Scala developers. +`finagle-stats` was written quite a bit later in order to +improve upon its implementation in terms of allocation efficiency and histogram precision. Why is `NullStatsReceiver` being used? @@ -34,21 +30,47 @@ and `Maven Central `_ Why migrate to finagle-stats? ----------------------------- -The `finagle-stats` library is designed to replace the other two libraries -formerly supported by Twitter: Ostrich and Commons Stats. The comparison table below -illustrates the advantages of `finagle-stats` over the older libraries. - -+------------------------------+---------------+----------+-----------------+ -| | Commons Stats | Ostrich | finagle-stats | -+==============================+===============+==========+=================+ -| Allocations per `Stat.add` | 16 bytes | 0 bytes | 0 bytes | -+------------------------------+---------------+----------+-----------------+ -| Highest histogram resolution | p99 | p9999 | p9999 | -+------------------------------+---------------+----------+-----------------+ -| Sampling error | N/A | ±5% | ±0.5% | -+------------------------------+---------------+----------+-----------------+ +The `finagle-stats` library is designed to replace Ostrich which was +formerly supported by Twitter. The comparison table below +illustrates the advantages of `finagle-stats` over Ostrich. + ++------------------------------+----------+-----------------+ +| | Ostrich | finagle-stats | ++==============================+==========+=================+ +| Allocations per `Stat.add` | 0 bytes | 0 bytes | ++------------------------------+----------+-----------------+ +| Highest histogram resolution | p9999 | p9999 | ++------------------------------+----------+-----------------+ +| Sampling error | ±5% | ±0.5% | ++------------------------------+----------+-----------------+ `finagle-stats` allows you to visualize and download the full details of `Stat` histograms. This is available on the `/admin/histogram `_ admin page. + +Are metric values reset or latched after an observability collection? +--------------------------------------------------------------------- +``Stats`` are latched by all two implementations. +This means that for each observability collection window, +you will see that window’s distribution of values. + +There is no notion of latching/resetting a ``Gauge``, as it is +always a discrete point in time value. This is the case across +all three implementations, except for CounterishGauges, which +creates a gauge that models an unlatched counter. + +``Counters`` have different behavior depending on which implementation +is being used. This in turn has an impact in how you write your +queries for your dashboards and alerts. Both Commons Metrics :sup:`[1]` and Commons Stats +are **not** “latched”, which means that their values are monotonically +increasing (assuming you never call ``incr()`` with a negative value). +This implies that you use `CQLs rate function `_ +if you want to see per minute counts. Ostrich’s ``Counters`` are latched +when they are collected (think of this as being a “reset”) +which means that you should not need to use ``rate`` in order +to see per minute values. + +:sup:`[1]` Note that by using ``finagle-stats`` and setting the flag +``com.twitter.finagle.stats.useCounterDeltas=true`` you can +get per-minute deltas in Commons Metrics. diff --git a/doc/src/sphinx/util-stats/user_guide.rst b/doc/src/sphinx/util-stats/user_guide.rst index 667d29949b..48cfcef612 100644 --- a/doc/src/sphinx/util-stats/user_guide.rst +++ b/doc/src/sphinx/util-stats/user_guide.rst @@ -88,6 +88,31 @@ along with code that holds a strong reference to the gauge in a global linked li such you should only rely on ``provideGauge`` when you do not have a place to keep a strong reference. + +Prefer fully-qualified names over scoping +----------------------------------------- +There is a convenient API method (``scope``) allowing to spawn a "scoped" version of a given +``StatsReceiver``. Although it enables effortless code reuse, scoping comes at the cost of extra +allocations. Inlining a fully-qualified metric name directly into the constructing method yields +the same outcome yet avoids the overhead needed to accommodate interim structures. + +Put this way, if possible, prefer this + +.. code-block:: scala + + statsReceiver.counter("foo", "bar", "baz") + +over this + +.. code-block:: scala + + statsReceiver.scope("foo").scope("bar").counter("baz") + +It's important to note that while this optimization is appealing like an easy win, it shouldn't +be universally applied. There are many legitimate use-cases when scoping a ``StatsReceiver`` +is very reasonable thing to do, and otherwise, would require an alternative channel for the scope +to be propagated between components, presumably introducing the overhead somewhere else. + .. _testing_code: Testing code that use StatsReceivers @@ -112,7 +137,6 @@ function **must also** be thread-safe. Leveraging Verbosity Levels --------------------------- - Introducing a new application metric (i.e., a histogram or a counter) is always a trade-off between its operational value and its cost within observability. Verbosity levels for StatsReceivers are aiming to reduce the observability cost of Finagle-based services by allowing for a granular control @@ -124,9 +148,22 @@ operations. Limiting the number of exported metrics via verbosity levels can reduce applications' operational cost. However taking this to extremes may drastically affect operability of your service. We -recommend using your judgment to make sure blacklisting a given metric will not reduce a process' +recommend using your judgment to make sure denylisting a given metric will not reduce a process' visibility. +Custom Histogram Percentiles +---------------------------- +If you need a custom histogram percentile, you can use the `MetricBuilder` interface, and populate +it via `MetricBuilder#withPercentiles`. + +.. code-block:: scala + + import com.twitter.finagle.stats.{DefaultStatsReceiver, StatsReceiver, MetricBuilder} + + val sr: StatsReceiver = DefaultStatsReceiver + val mb: MetricBuilder = sr.metricBuilder().withPercentiles(Seq(0.99, 0.999, 0.9999, 0.90, 0.6)) + val stat = mb.histogram("my", "cool", "histo") + Access needed to a StatsReceiver in an inconvenient place --------------------------------------------------------- Ideally classes would be passed a properly scoped ``StatsReceiver`` in their constructor but this @@ -145,7 +182,7 @@ developers may find useful. - `NullStatsReceiver`_ for when you do not care about all metrics. -- `BlacklistStatsReceiver`_ programmatically decide which metrics to ignore. +- `DenylistStatsReceiver`_ programmatically decide which metrics to ignore. - `BroadcastStatsReceiver`_ allows for sending metrics to two or more ``StatsReceiver``\s. @@ -157,14 +194,14 @@ using. Via TwitterServer/finagle-stats — the `HTTP admin interface`_ responds with json at ``/admin/metrics.json`` and there is a web UI for watching them in real-time at ``/admin/metrics``. -.. _LoadService: https://github.com/twitter/finagle/blob/master/finagle-core/src/main/scala/com/twitter/finagle/util/LoadService.scala +.. _LoadService: https://github.com/twitter/finagle/blob/release/finagle-core/src/main/scala/com/twitter/finagle/util/LoadService.scala .. _finagle-stats: https://github.com/twitter/finagle/tree/master/finagle-stats -.. _BroadcastStatsReceiver: https://github.com/twitter/util/blob/master/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala +.. _BroadcastStatsReceiver: https://github.com/twitter/util/blob/release/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala .. _NullStatsReceiver: https://github.com/twitter/util/blob/develop/util-stats/src/main/scala/com/twitter/finagle/stats/NullStatsReceiver.scala .. _Futures: https://twitter.github.io/finagle/guide/Futures.html .. _Stat.time: https://github.com/twitter/util/blob/develop/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala .. _Stat.timeFuture: https://github.com/twitter/util/blob/develop/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala -.. _java.lang.ref.WeakReference: http://docs.oracle.com/javase/8/docs/api/java/lang/ref/WeakReference.html -.. _InMemoryStatsReceiver: https://github.com/twitter/util/blob/master/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala +.. _java.lang.ref.WeakReference: https://docs.oracle.com/javase/8/docs/api/java/lang/ref/WeakReference.html +.. _InMemoryStatsReceiver: https://github.com/twitter/util/blob/release/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala .. _HTTP admin interface: https://twitter.github.io/twitter-server/Features.html#http-admin-interface -.. _BlacklistStatsReceiver: https://github.com/twitter/util/blob/develop/util-stats/src/main/scala/com/twitter/finagle/stats/BlacklistStatsReceiver.scala +.. _DenylistStatsReceiver: https://github.com/twitter/util/blob/develop/util-stats/src/main/scala/com/twitter/finagle/stats/DenylistStatsReceiver.scala diff --git a/doc/src/sphinx/util-validator/index.rst b/doc/src/sphinx/util-validator/index.rst new file mode 100644 index 0000000000..6b497aae11 --- /dev/null +++ b/doc/src/sphinx/util-validator/index.rst @@ -0,0 +1,1820 @@ +.. _util-validator-index: + +util-validator Guide +==================== + +Preface +------- + +Validating data is a common task that occurs throughout all application layers, from the +presentation to the persistence layer. Often the same validation logic is implemented in each +layer which is time consuming and error-prone. To avoid duplication of these validations, +developers often bundle validation logic directly into the domain model, cluttering domain classes +with validation code which is really metadata about the class itself. + +The `Jakarta Bean Validation 2.0 `__ - defines a metadata model and API +for entity and method validation. The default metadata source are Java annotations (with the ability +to override and extend the meta-data through the use of XML if desired). The API is not tied to a +specific application tier nor programming model and it is specifically not tied to either the web or +persistence tier, and is available for both server-side application programming or rich client +applications. + +The `ScalaValidator` is a Scala wrapper around the Java `Hibernate Validator `__ +which itself is the reference implementation of `Jakarta Bean Validation specification `__. + +Getting Started +--------------- + +The `ScalaValidator` depends only on the `Hibernate Validator `__, +the `Glassfish reference implementation `__ of the +`Jakarta Expression Language `__ for message interpolation, +and `json4s-core `__ for Scala reflection. + +We highly recommend reading through the Hibernate Validator `documentation `__ +as our guide is largely based on this information. + +In order to use the `util-validator` add a dependency on the project, e.g. in `sbt `__: + +.. code-block:: scala + + "com.twitter" %% "util-validator" % utilValidatorVersion + +or with `Maven `__: + +.. code-block:: xml + + + com.twitter + util-validator_2.12 + ${util.validator.version} + + +where `utilValidatorVersion` and `util.validator.version` are the expected library version. + +Applying constraints +~~~~~~~~~~~~~~~~~~~~ + +Let's start with a basic example to see how to apply constraints. Let's assume you have a +case class defined as such: + +.. code-block:: scala + + package com.example + + import jakarta.validation.constraints.{Min, NotEmpty, Size} + + case class Car( + @NotEmpty manufacturer: String, + @NotEmpty @Size(min = 2, max = 14) licensePlate: String, + @Min(2) seatCount: Int) + +The `@NotEmpty`, `@Size` and `@Min` annotations are used to declare the constraints which should be +applied to the fields of a `Car` instance: + +* `manufacturer` must never be an empty String (and thus also not `null`) +* `licensePlate` must never be an empty String (thus not also not `null`) and must be between 2 and 14 characters +* `seatCount` must be equal to at least 2 + +Validating constraints +~~~~~~~~~~~~~~~~~~~~~~ + +`ScalaValidator#validate` +^^^^^^^^^^^^^^^^^^^^^^^^^ + +To perform a validation of these constraints, you use an instance of the `ScalaValidator`. Let’s +take a look at a test for validating `Car` instances: + +.. code-block:: scala + :emphasize-lines: 15,25,35,45 + + package com.example + + import com.twitter.util.validation.ScalaValidator + import jakarta.validation.ConstraintViolation + import org.scalatest.funsuite.AnyFunSuite + import org.scalatest.matchers.should.Matchers + + class CarTest extends AnyFunSuite with Matchers { + + private[this] val validator: ScalaValidator = ScalaValidator() + + test("manufacturer is empty") { + val car: Car = Car("", "DD-AB-123", 4) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should equal("manufacturer") + violation.getMessage should be("must not be empty") + } + + test("licensePlate is too short") { + val car: Car = Car("Greenwich", "D", 4) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should equal("licensePlate") + violation.getMessage should be("size must be between 2 and 14") + } + + test("seatCount is too small") { + val car: Car = Car("Greenwich", "DD-AB-123", 1) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should equal("seatCount") + violation.getMessage should be("must be greater than or equal to 2") + } + + test("car is valid") { + val car: Car = Car("Greenwich", "DD-AB-123", 2) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + violations.isEmpty should be(true) + } + } + +In the test constructor we instantiate an instance of a `ScalaValidator`. A `ScalaValidator` instance +is thread-safe and may be reused multiple times. It thus can safely be stored in a member variable +field and be used in the test cases to validate the different `Car` instances. + +The `ScalaValidator#validate` method returns a set of `ConstraintViolation` instances, which you can +iterate over in order to see which validation errors occurred. The first three tests show some +expected constraint violations: + +* The `@NotEmpty` constraint on `manufacturer` is violated in `test("manufacturer is empty")` +* The `@Size` constraint on `licensePlate` is violated in `test("licensePlate is too short")` +* The `@Min` constraint on `seatCount` is violated in `test("seatCount is too small")` + +If the object validates successfully, `ScalaValidator#validate` returns an empty set as you can see +in `test("car is valid")`. + +Note that this method is recursive and will cascade validations as explained in the section on +`Object graphs <#object-graphs>`__. + +`ScalaValidator#verify` +^^^^^^^^^^^^^^^^^^^^^^^ + +The `ScalaValidator` also has an API to throw constraint violations as a type of `ValidationException` +rather than returning a `Set[ConstraintViolation[_]`. + +Instead of calling `ScalaValidator#validate`, use `ScalaValidator#verify`: + +.. code-block:: scala + :emphasize-lines: 16,28,40,52 + + package com.example + + import com.twitter.util.validation.ScalaValidator + import jakarta.validation.{ConstraintViolation, ConstraintViolationException} + import org.scalatest.funsuite.AnyFunSuite + import org.scalatest.matchers.should.Matchers + + class CarTest extends AnyFunSuite with Matchers { + + private[this] val validator: ScalaValidator = ScalaValidator() + + test("manufacturer is empty") { + val car: Car = Car("", "DD-AB-123", 4) + + val e = intercept[ConstraintViolationException] { + validator.verify(car) + } + e.getConstraintViolations.size should equal(1) + val violation = e.getConstraintViolations.iterator.next + violation.getPropertyPath.toString should equal("manufacturer") + violation.getMessage should be("must not be empty") + } + + test("licensePlate is too short") { + val car: Car = Car("Greenwich", "D", 4) + + val e = intercept[ConstraintViolationException] { + validator.verify(car) + } + e.getConstraintViolations.size should equal(1) + val violation = e.getConstraintViolations.iterator.next + violation.getPropertyPath.toString should equal("licensePlate") + violation.getMessage should be("size must be between 2 and 14") + } + + test("seatCount is too small") { + val car: Car = Car("Greenwich", "DD-AB-123", 1) + + val e = intercept[ViolationException] { + validator.verify(car) + } + e.isInstanceOf[ConstraintViolationException] should be(true) + e.asInstanceOf[ConstraintViolationException].getConstraintViolations.size should equal(1) + val violation = e.asInstanceOf[ConstraintViolationException].getConstraintViolations.iterator.next + violation.getPropertyPath.toString should equal("seatCount") + violation.getMessage should be("must be greater than or equal to 2") + } + + test("car is valid") { + val car: Car = Car("Greenwich", "DD-AB-123", 2) + + validator.verify(car) + } + } + +Like `ScalaValidator#validate`, this method is recursive and will cascade validations as explained +in the section on `Object graphs <#object-graphs>`__. + +Declaring and validating case class constraints +----------------------------------------------- + +In this section we show how to declare (see `“Declaring case class constraints” <#declaring-case-class-constraints>`__) +and validate (see `“Validating bean constraints” <#validating-case-class-constraints>`__) case class +constraints. + +You will likely want to review Hibernate's `“Built-in constraints” `__ +documentation which provides an overview of all built-in constraints coming with the Hibernate +Validator and thus supported here. + +If you are interested in applying constraints to method parameters and return values, refer to +`Declaring and validating method constraints <#declaring-and-validating-method-constraints>`__. + +Declaring case class constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Constraints in Jakarta Bean Validation are expressed via Java annotations. In this section we show +how to enhance an object model with these annotations. There are three types of bean constraints: + +* field constraints +* property constraints +* container element constraints +* class constraints + +When defining a class field in Scala the compiler creates up to four accessors for it: a getter, +a setter, and if the field is annotated with `@BeanProperty`, a bean getter and a bean setter +[`reference `__]. + +Thus, if you had the following class definition: + +.. code-block:: scala + + class C(@myAnnot var c: Int) + +There are **six** entities which can carry the `@myAnnot` annotation: the constructor param, the +generated field and the four accessors. **By default, annotations on constructor parameters end up +only on the constructor parameter and not on any other entity**. Annotations on fields, by default, +only end up on the field. + +If you wanted Scala to copy the annotation to the generated field, you would need to apply the +`scala.annotation.meta.field` meta annotation to the annotation. + +For example, using `class C` again from above: + +.. code-block:: scala + + class C(@(myAnnot @field) var c: Int) + +Thus to use the Hibernate Validator Java library directly from Scala, you would need to always +annotate constraint annotations with a `scala.annotation.meta` annotation along with translating +between Scala and Java collection types. + +The benefit of `util-validator` is that the library will work with simply annotated case class +constructor parameters (i.e., without needing to add a `scala.annotation.meta` annotation) with an +API that uses and expresses Scala collection types. + +Constructor param constraints +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +Constraints can be expressed by annotating the constructor param of a case class. E.g., + +.. code-block:: scala + + package com.example + + import jakarta.validation.constraints.{Min, NotEmpty, Size} + + case class Car( + @NotEmpty manufacturer: String, + @NotEmpty @Size(min = 2, max = 14) licensePlate: String, + @Min(2) seatCount: Int) + +.. important:: + + When using constructor parameter constraints, field access strategy is used to access the value + to be validated. This means the validation logic directly accesses the instance variable and + will never invoke the generated property accessor method. + +Constraints can be applied to fields of any access type (public, private etc.), including inherited +members. Neither constraints on companion object fields, nor secondary companion object `#apply` +method constraints are supported. + +Secondary constructors +++++++++++++++++++++++ + +The validation library is not used for case class construction but because of how annotations are +carried by default in Scala for classes, the constructor of a case class plays an important role +in case class validation. By default, the library only tracks constraint annotations for constructor +parameters which are also declared fields within the class. Thus, if you have a primary constructor +of `Int` types with a secondary constructor of constraint-annotated `String` types (which does some +translation into the required `Int` types), the secondary constructor constraints have no bearing on +validation. + +Container element constraints +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +Note that due to a long-standing issue [`SCALA-9883 `__] +which affects how to access Java annotations from the Scala compiler, that **container element** +**constraints are not supported**. + +For examples of container element constraints, see the Hibernate `documentation `__. + +Thus, take note that constraint annotations will apply only to the container and not the type(s) of +the container with the one caveat being a field of type `Option[_]` which is always unwrapped by +default and thus the annotation always pertains to the contained type. + +With `Option[_]` +++++++++++++++++ + +When applying a constraint on the type argument of `Option[_]`, the `ScalaValidator` will +*automatically unwrap* the `Option[_]` type and validate the internal value. For example: + +.. code-block:: bash + :emphasize-lines: 21,22 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import jakarta.validation.constraints._ + import jakarta.validation.constraints._ + + scala> case class Car(@Min(1000) towingCapacity: Option[Int] = None) + defined class Car + + scala> val car = Car(Some(100)) + car: Car = Car(Some(100)) + + scala> import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.ScalaValidator + + scala> val validator = ScalaValidator() + Mar 30, 2021 11:58:47 AM org.hibernate.validator.internal.util.Version + INFO: HV000001: Hibernate Validator 7.0.1.Final + validator: com.twitter.util.validation.ScalaValidator = com.twitter.util.validation.ScalaValidator@574f9e36 + + scala> val violations = validator.validate(car) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must be greater than or equal to 1000', propertyPath=towingCapacity, rootBeanClass=class Car, messageTemplate='{jakarta.validation.constraints.Min.message}'}) + + scala> val violation = violations.head + violation: jakarta.validation.ConstraintViolation[Car] = ConstraintViolationImpl{interpolatedMessage='must be greater than or equal to 1000', propertyPath=towingCapacity, rootBeanClass=class Car, messageTemplate='{jakarta.validation.constraints.Min.message}'} + + scala> violation.getMessage + res0: String = must be greater than or equal to 1000 + + scala> violation.getPropertyPath.toString + res1: String = towingCapacity + + scala> + +.. note:: + + The property path only contains the name of the property, in this way, `Options` are treated + as a "transparent" container. + +Class-level constraints +^^^^^^^^^^^^^^^^^^^^^^^ + +A constraint can also be placed on the class level. In this case not a single property is subject +of the validation but the complete object. Class-level constraints are useful if the validation +depends on a correlation between several properties of an object. + +For instance, if we had a `Car` case class defined with two field `seatCount` and `passengers` and +we want to ensure that the list of `passengers` does not have more entries than available seats. +For that purpose the `@ValidPassengerCount` constraint is added on the class level. The validator of +that constraint has access to the complete `Car` object, allowing to compare the numbers of seats +and passengers. + +.. code-block:: scala + :emphasize-lines: 5 + + package com.example + + import jakarta.validation.constraints.Min + + @ValidPassengerCount + case class Car( + @Min(2) seatCount: Int, + passengers: Seq[Person]) + +See the corresponding section in `Creating custom constraints - Class-level constraints <#>`__ for +documentation on how to implement such a custom constraint. + +Constraint inheritance +^^^^^^^^^^^^^^^^^^^^^^ + +When a class implements a trait or extends an abstract class, all constraint annotations declared +on the super-type apply in the same manner as the constraints specified in the class itself. +For example, if we had a trait, `Car` with a `RentalCar` implementation: + +.. code-block:: scala + + trait Car { + @NotEmpty def manufacturer: String + } + + case class RentalCar(manufacturer: String, @NotEmpty rentalStation: String) extends Car + +Here the case class `RentalCar` implements the `Car` trait and adds the property `rentalStation`. If +an instance of `RentalCar` is validated, not only the `@NotEmpty` constraint on `rentalStation` is +evaluated, but also the constraint on `manufacturer` from the `Car` trait. + +.. code-block:: bash + :emphasize-lines: 26,27,41,42,56,57,70,71 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import jakarta.validation.constraints._ + import jakarta.validation.constraints._ + + scala> trait Car { + | @NotEmpty def manufacturer: String + | } + defined trait Car + + scala> case class RentalCar(manufacturer: String, @NotEmpty rentalStation: String) extends Car + defined class RentalCar + + scala> import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.ScalaValidator + + scala> val validator = ScalaValidator() + Mar 30, 2021 12:21:41 PM org.hibernate.validator.internal.util.Version + INFO: HV000001: Hibernate Validator 7.0.1.Final + validator: com.twitter.util.validation.ScalaValidator = com.twitter.util.validation.ScalaValidator@64e89bb2 + + scala> val rental = RentalCar("", "Hertz") + rental: RentalCar = RentalCar(,Hertz) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=manufacturer, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.size + res0: Int = 1 + + scala> violations.head.getPropertyPath.toString + res: String = manufacturer + + scala> violations.head.getMessage + res: String = must not be empty + + scala> val rental = RentalCar("Renault", "") + rental: RentalCar = RentalCar(Renault,) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=rentalStation, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.size + res: Int = 1 + + scala> violations.head.getPropertyPath.toString + res: String = rentalStation + + scala> violations.head.getMessage + res: String = must not be empty + + scala> val rental = RentalCar("", "") + rental: RentalCar = RentalCar(,) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=manufacturer, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}, ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=rentalStation, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.size + res: Int = 2 + + scala> violations.map(v => s"${v.getPropertyPath}: ${v.getMessage}").mkString("\n") + res: String = + manufacturer: must not be empty + rentalStation: must not be empty + + scala> val rental = RentalCar("Renault", "Hertz") + rental: RentalCar = RentalCar(Renault,Hertz) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set() + + scala> + +The same would be true, if `Car` was not a trait but an abstract class implemented by `RentalCar`. + +**Constraint annotations are aggregated if members are overridden**. So if the `RentalCar` case class +also specifies any constraints on the `manufacturer` field, it would be evaluated in addition to the +`@NotEmpty` constraint from the `Car` trait: + +.. code-block:: bash + :emphasize-lines: 7,8,22,23 + + scala> case class RentalCar(@Size(min = 2, max = 14) manufacturer: String, @NotEmpty rentalStation: String) extends Car + defined class RentalCar + + scala> val rental = RentalCar("A", "Hertz") + rental: RentalCar = RentalCar(A,Hertz) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='size must be between 2 and 14', propertyPath=manufacturer, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.Size.message}'}) + + scala> violations.size + res: Int = 1 + + scala> violations.head.getPropertyPath.toString + res: String = manufacturer + + scala> violations.head.getMessage + res: String = size must be between 2 and 14 + + scala> val rental = RentalCar("", "Hertz") + rental: RentalCar = RentalCar(,Hertz) + + scala> val violations = validator.validate(rental) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='size must be between 2 and 14', propertyPath=manufacturer, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.Size.message}'}, ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=manufacturer, rootBeanClass=class RentalCar, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.size + res: Int = 2 + + scala> violations.map(v => s"${v.getPropertyPath}: ${v.getMessage}").mkString("\n") + res: String = + manufacturer: size must be between 2 and 14 + manufacturer: must not be empty + + scala> + +Object graphs +^^^^^^^^^^^^^ + +The Jakarta Bean Validation API does not only allow to validate single case class instances but also +complete object graphs via cascaded validation. To do so, just annotate a field or property +representing a reference to another case class with `@Valid`. + +.. code-block:: scala + + case class Person(@NotEmpty name: String) + + case class Car(@NotEmpty manufacturer: String, @Valid driver: Person) + +If an instance of `Car` is validated, the referenced `Person` case class will also be validated, as +the `driver` field is annotated with `@Valid`. Therefore the validation of a `Car` will fail if the +`name` field of the referenced Person instance is empty. + +The validation of object graphs is recursive, i.e. if a reference marked for cascaded validation +points to an object which itself has properties annotated with `@Valid`, these references will be +followed up by the validation logic as well. The validation logic will ensure that no infinite +loops occur during cascaded validation, for example if two objects hold references to each other. + +Note that `null` or `None` values are ignored during cascaded validation. + +`Iterable[_]` ++++++++++++++ + +Unlike constraints, object graph validation works for some container elements (those which are +an `Iterable[_]`). That means the field of the container can be annotated with `@Valid`, which will +cause each contained element to be validated when the parent object is validated. + +.. code-block:: scala + + case class Person(@NotEmpty name: String) + + case class Car(@NotEmpty manufacturer: String, @Valid drivers: Seq[Person]) + +If an instance of `Car` is validated, the referenced sequence of `Person` case classes will also +be validated, as the `drivers` field is annotated with `@Valid`. Therefore the validation of a `Car` +will fail if the `name` field of the referenced Person instance is empty. + +.. code-block:: bash + :emphasize-lines: 27,28,39,40 + + Welcome to Scala 2.12.13 (JDK 64-Bit Server VM, Java 1.8.0_242). + Type in expressions for evaluation. Or try :help. + + scala> import jakarta.validation.constraints._ + import jakarta.validation.constraints._ + + scala> case class Person(@NotEmpty name: String) + defined class Person + + scala> import jakarta.validation.Valid + import jakarta.validation.Valid + + scala> case class Car(@NotEmpty manufacturer: String, @Valid drivers: Seq[Person]) + defined class Car + + scala> val car = Car("Renault", Seq(Person(""))) + car: Car = Car(Renault,List(Person())) + + scala> import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.ScalaValidator + + scala> val validator = ScalaValidator() + Mar 30, 2021 3:59:53 PM org.hibernate.validator.internal.util.Version + INFO: HV000001: Hibernate Validator 7.0.1.Final + validator: com.twitter.util.validation.ScalaValidator = com.twitter.util.validation.ScalaValidator@6dcc7696 + + scala> val violations = validator.validate(car) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=drivers[0].name, rootBeanClass=class Car, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.head.getPropertyPath.toString + res: String = drivers[0].name + + scala> violations.head.getMessage + res: String = must not be empty + + scala> val car = Car("Renault", Seq(Person("Lupin"), Person(""))) + car: Car = Car(Renault,List(Person(Lupin), Person())) + + scala> val violations = validator.validate(car) + violations: Set[jakarta.validation.ConstraintViolation[Car]] = Set(ConstraintViolationImpl{interpolatedMessage='must not be empty', propertyPath=drivers[1].name, rootBeanClass=class Car, messageTemplate='{jakarta.validation.constraints.NotEmpty.message}'}) + + scala> violations.head.getPropertyPath.toString + res: String = drivers[1].name + + scala> violations.head.getMessage + res: String = must not be empty + + scala> + +Validating case class constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The `ScalaValidator` does not directly implement the Jakarta Bean Validation interface but provides +the same methods (and a few more) with Scala collection types. Note, however, since the +`ScalaValidator` is a wrapper that the underlying Jakarta Bean Validation is available for direct +access for Java users or to use functionality not exposed in the `ScalaValidator`. Also note that +this is a Java library and bypassing the `ScalaValidator` methods will thus not support all Scala +types. + +This section shows how to obtain a configured `ScalaValidator` instance and then will walk through +the different methods of the `ScalaValidator`. + +Obtaining a Validator instance +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The `ScalaValidator` has an `apply()` function which will return an instance of a `ScalaValidator` +configured with defaults. + +.. code-block:: scala + + val validator: ScalaValidator = ScalaValidator() + +To apply custom configuration to create a `ScalaValidator`, there is a builder for constructing a +customized validator. + +E.g., to configure a custom size for the descriptor cache: + +.. code-block:: scala + :emphasize-lines: 3 + + val validator: ScalaValidator = + ScalaValidator.builder + .withDescriptorCacheSize(1024) + .validator + +Or to add custom validations: + +.. code-block:: scala + :emphasize-lines: 10 + + import com.twitter.util.validation.cfg.ConstraintMapping + + val mapping = ConstraintMapping( + annotationType = classOf[FooConstraintAnnotation], + constraintValidator = classOf[FooConstraintValidator] + ) + + val validator: ScalaValidator = + ScalaValidator.builder + .withConstraintMapping(mapping) + .validator + +or to specify multiple constraint mappings: + +.. code-block:: scala + :emphasize-lines: 17 + + import com.twitter.util.validation.cfg.ConstraintMapping + + val mappings: Set[ConstraintMapping] = + Set( + ConstraintMapping( + annotationType = classOf[FooConstraintAnnotation], + constraintValidator = classOf[FooConstraintValidator] + ), + ConstraintMapping( + annotationType = classOf[BarConstraintAnnotation], + constraintValidator = classOf[BarConstraintValidator] + ) + ) + + val validator: ScalaValidator = + ScalaValidator.builder + .withConstraintMappings(mappings) + .validator + +You are **not required** to add a constraint mapping if the constraint annotation specifies a +`validatedBy()` class. + +.. important:: + + Attempting to add multiple mappings for the same annotation will result in a `ValidationException `__. + +ScalaValidator API +^^^^^^^^^^^^^^^^^^ + +The `ScalaValidator` contains methods that can be used to either validate entire entities or just +single properties of the entity. All methods but `ScalaValidator#verify` return a +`Set[ConstraintViolation[_]]`. The set is empty if the validation succeeds. Otherwise, a +`ConstraintViolation` instance is added for each violated constraint. In the case of +`ScalaValidator#verify`, when the returned `Set[ConstraintViolation[_]]` is non-empty a +`ConstraintViolationException `__ +is thrown which wraps the `Set[ConstraintViolation[_]]`. + +All the validation methods have a version which takes a `groups: Seq[Class[_]]` parameter that is +used to specify validation groups to be considered when performing the validation. If the parameter +is empty, the default validation group (`jakarta.validation.groups.Default`) is used. + +The topic of validation groups is discussed in detail in Hibernate's `Chapter 5, Grouping constraints `__. + +`ScalaValidator#validate` ++++++++++++++++++++++++++ + +Use `validate` to perform validation of all constraints of a given case class. + +.. code-block:: scala + :emphasize-lines: 8 + + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{AssertTrue, NotEmpty} + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) + + val car = Car("", false) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + assert(violations.size == 2) + assert(violations.head.getMessage == "must not be empty") + assert(violations.last.getMessage == "not true") + +Using the given definition of a `Car`, the above example shows an instance which fails to satisfy +the `@NotEmpty` constraint on the `manufacturer` property. The validation call therefore returns +one `ConstraintViolation[Car]` object. + +`ScalaValidator#verify` ++++++++++++++++++++++++ + +The `verify` method is the same as `validate` but instead of returning `Set[ConstraintViolation[T]]` +the method throws a `ConstraintViolationException `__. + +.. code-block:: scala + :emphasize-lines: 9 + + import com.twitter.util.Try + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{AssertTrue, NotEmpty} + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) + + val car = Car("", false) + + val result = Try(validator.verify(car)) + assert(result.isThrow) + +`ScalaValidator#validateValue` +++++++++++++++++++++++++++++++ + +By using the `validateValue` method you can check whether a single field of a given case class can +be validated successfully if the field had the given value: + +.. code-block:: scala + :emphasize-lines: 7,8,9,10 + + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{AssertTrue, NotEmpty} + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) + + val violations: Set[ConstraintViolation[Car]] = + validator.validateValue( + classOf[Car], + "manufacturer", + "") + assert(violations.size == 1) + assert(violations.head.getMessage == "must not be empty") + +There is also a version of the method which instead of accepting a `Class[T]`, takes a +`CaseClassDescriptor[T]` of a given class. + +.. code-block:: scala + :emphasize-lines: 6,9,10,11,12 + + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{AssertTrue, NotEmpty} + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) + + val carDescriptor: CaseClassDescriptor[Car] = validator.getConstraintsForClass(classOf[Car]) + + val violations: Set[ConstraintViolation[Car]] = + validator.validateValue( + carDescriptor, + "manufacturer", + "") + assert(violations.size == 1) + assert(violations.head.getMessage == "must not be empty") + +`ScalaValidator#validateProperty` ++++++++++++++++++++++++++++++++++ + +With help of the `validateProperty` you can validate a single named field of a given case class +instance. + +.. code :: scala + :emphasize-lines: 9,10,11 + + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{AssertTrue, NotEmpty} + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) + + val car = Car("", false) // isRegistered is false here which would normally fail `@AssertTrue` validation + + val violations: Set[ConstraintViolation[Car]] = + validator.validateProperty( + car, + "manufacturer") + assert(violations.size == 1) + assert(violations.head.getMessage == "must not be empty") + +.. note:: + + The difference between `validateValue` and `validateProperty` is that the former takes a class + type, a field in the class, and a presumed value in order to perform validation. Thus, no actual + instance of the class is required for performing validation. + + Whereas the latter performs validation of a given field on an instance of a case class. + +.. important:: + + The `@Valid` annotation for cascading validation is not honored by `validateValue` or + `validateProperty`. + +ConstraintViolation +^^^^^^^^^^^^^^^^^^^ + +As mentioned, all methods but `ScalaValidator#verify` return a `Set[ConstraintViolation[_]]`. The +`ConstraintViolation `__ +contains information about the cause and location of the validation failure. + +See the `Hibernate documentation `__ for details. + +Built-in constraints +^^^^^^^^^^^^^^^^^^^^ + +In addition to the constraints defined by the `Jakarta Bean Validation API and the Hibernate Built-in constraints `__ +we provide several useful custom constraints which are listed below. These constraints all apply to +the field/property level. + +* `@CountryCode` + + Checks that the annotated field represents a valid `ISO 3166 `__ country code or a collection of `ISO 3166 `__ country codes. + + - **Supported data types** + + `String`, `Array[String]`, `Iterable[String]` + +* `@OneOf(value=)` + + Checks that the annotated field represents a value or a set of values from the given array of values. Note that the `#toString` representations are checked for equality. + + - **Supported data types** + + `Any`, `Array[Any]`, `Iterable[Any]` + +* `@UUID` + + Checks if the annotated field is a valid `String` representation of a `java.util.UUID `__. + + - **Supported data types** + + `String` + +Declaring and validating method constraints +------------------------------------------- + +Constraints can not only be applied to case classes and their fields, but also to the parameters and +return values of methods within the case class. Following the Jakarta Bean Validation specification, +constraints can be used to specify + +* the preconditions that must be satisfied by the caller before a method or constructor may be invoked (by applying constraints to the parameters of an executable) +* the postconditions that are guaranteed to the caller after a method or constructor invocation returns (by applying constraints to the return value of an executable) + +.. important:: + + In all cases **except when using** `@MethodValidation`, declaring method or constructor + constraints itself does not automatically cause their validation upon validation of a case class + instance. Instead, the `ScalaExecutableValidator` API must be used to perform the validation. + +Declaring method constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +`@MethodValidation` constraint +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The `@MethodValidation` constraint annotation is a special method validation constraint provided by +the library. + +Unlike other executable specific constraints, `@MethodValidation`-annotated methods **will be** +**validated along with all field-level and class-level validations** when calling +`ScalaValidate#validate <#id1>`__ or `ScalaValidator#verify <#id2>`__. Note, that the method +validation is strictly performed *after* all field level validations (and *before* class-level +validations). Annotated methods MUST have no specified parameters and MUST return a +`MethodValidationResult`. + +To manually validate a given method or methods, you can use `ScalaExecutableValidator#validateMethod <#scalaexecutablevalidator-validatemethod>`__ +or `ScalaExecutableValidator#validateMethods <#scalaexecutablevalidator-validatemethods>`__. + +You can find details in `“Creating custom constraints - @MethodValidation constraints” <#methodvalidation-constraints>`__ +which shows how to implement a `@MethodValidation` constraint. + +Method Parameter constraints +^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +You can specify the preconditions of a method by adding constraint annotations to its parameters: + +.. code-block:: scala + + import jakarta.validation.constraints.{Min, NotEmpty, NotNull} + import java.time.LocalDate + + case class RentalStation(@NotEmpty name: String) { + + def rentCar( + @NotNull customer: Customer, + @NotNull @Future startDate: LocalDate, + @Min(1) durationInDays: Int + ): Unit = ??? + } + +The following preconditions are specified here: + +* The `name` passed to the `RentalCar` constructor must not be empty. +* When invoking the `rentCar()` method, the given customer must not be null, the rental’s start date must not be `null` as well as be in the future and finally the rental duration must be at least one day. + +.. note:: + + Constraints may only be applied to instance methods, i.e. declaring constraints on static + methods is not supported. + +Cross-parameter constraints ++++++++++++++++++++++++++++ + +Sometimes validation does not only depend on a single parameter but on several or even all +parameters of a method. This kind of requirement can be fulfilled with help of a cross-parameter +constraint. + +Cross-parameter constraints can be considered as the method validation equivalent to class-level +constraints. Both can be used to implement validation requirements which are based on several +elements. While class-level constraints can apply to several fields of a case class, cross-parameter +constraints can apply to several parameters of a method. + +In contrast to single-parameter constraints, cross-parameter constraints are declared on the method. +In the example below, the cross-parameter constraint `@LuggageCountMatchesPassengerCount` declared +on the `load()` method is used to ensure that no passenger has more than two pieces of luggage. + +.. code-block:: scala + + case class Car(@NotEmpty manufacturer: String, @AssertTrue isRegistered: Boolean) { + + @LuggageCountMatchesPassengerCount(piecesOfLuggagePerPassenger = 2) + def load(passengers: Seq[Person], luggage: Seq[PieceOfLuggage]): Unit = ??? + } + +We'll see in the next section that return value constraints are also declared on the method level. +Thus, in order to distinguish cross-parameter constraints from return value constraints, the +constraint target is configured in the `ConstraintValidator` implementation using the `@SupportedValidationTarget `__ +annotation. + +More details in `“Creating custom constraints - Cross-parameter constraints” <#id4>`__ which shows +how to implement a cross-parameter constraint. + +Method Return value constraints +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The postconditions of a method or constructor are declared by adding constraint annotations to the +executable: + +.. code-block:: scala + + case class RentalStation @ValidRentalStation(@NotEmpty name: String) { + + @NotEmpty + @Size(min = 1) + def getCustomers: Seq[Customer] = ??? + } + +In the example above, the following constraints apply to executables (constructor and a method in +this case) of `RentalStation`: +* A new `RentalStation` instance must satisfy the `@ValidRentalStation` constraint +* The `Seq[Customer]` returned by `getCustomers` must not be empty and must contain at least on element + +Cascaded Validation +^^^^^^^^^^^^^^^^^^^ + +Like the cascaded validation of case class fields (see `“Object graphs” <#object-graphs>`__), the +`@Valid` annotation can be used to mark executable parameters and return values for cascaded +validation. When validating a parameter or return value annotated with `@Valid`, the constraints +declared on the parameter or return value object are also validated. + +.. code-block:: scala + + case class Car(@NotEmpty manufacturer: String, @NotEmpty @Size(min = 2, max = 14) licensePlate: String) + + case class Garage @Valid (@NotEmpty name: String) { + + def checkCar(@NotNull @Valid car: Car): Boolean = ??? + } + +Above, the `car` parameter of the method `Garage#checkCar` as well as the return value of the +`Garage` constructor are marked for cascaded validation. + +When validating the arguments of the `checkCar()` method, the constraints on the properties of the +passed `Car` object are evaluated as well. Similarly, the `@NotEmpty` constraint on the name field +of `Garage` is checked when validating the return value of the `Garage` constructor. + +Generally, the cascaded validation works for executables in the same way as it does for case class +fields. + +In particular, `null` or `None` values are ignored during cascaded validation (which is impossible +for constructor return values) and cascaded validation is performed recursively, i.e. if a parameter +or return value object which is marked for cascaded validation itself has properties marked with +`@Valid`, the constraints declared on the referenced elements will be validated as well. + +Also, the same issues with annotating container elements apply because of the Scala compiler +limitations but also note that more generally, the majority of executable methods proxy to the +underlying Hibernate Java library. + +Method constraints in inheritance hierarchies +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +Note there are some rules for method constraints and inheritance as outlined in the +`Hibernate documentation `__. + +Validating method constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The validation of method constraints is done using the `ScalaExecutableValidator` interface. + +Obtaining a `ScalaExecutableValidator` instance +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The `ScalaValidator` provides a `ScalaExecutableValidator` which implements a set of "executable" +methods that mirrors the bean validation `ExecutableValidator `__ API. + +All of the corollary `ExecutableValidator` methods proxy directly to the underlying Hibernate +Validator for support since it properly handles annotations in these cases and no special logic is +needed. + +The `ScalaExecutableValidator` can be obtained from the `ScalaValidator` by calling +`ScalaValidator#forExecutables` + +.. code-block:: scala + + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + +`ScalaExecutableValidator` methods +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The `ScalaExecutableValidator` interface offers altogether six methods: + +* `validateMethods()`, `validateMethod()`, `validateParameters()`, and `validateReturnValue()` for method validation +* `validateConstructorParameters()` and `validateConstructorReturnValue()` for constructor validation + +`ScalaExecutableValidator#validateMethods` +++++++++++++++++++++++++++++++++++++++++++ + +Validates all `@MethodValidation`-annotated methods of a given object. + +.. code-block:: scala + :emphasize-lines: 27 + + import com.twitter.util.validation.MethodValidation + import com.twitter.util.validation.engine.MethodValidationResult + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{Min, NotEmpty} + import java.time.LocalDate + + case class RentalCar( + @NotEmpty make: String, + @NotEmpty model: String, + @Min(2000) modelYear: Int) { + + @MethodValidation(fields = Array("modelYear")) + def onlyNewerCars: MethodValidationResult = { + // only want model years within the last 2 years + val year: Int = LocalDate.now.getYear + if ((year - modelYear) <= 2) MethodValidationResult.Valid() + else MethodValidationResult.Invalid("model year must be within the last 2 years") + } + } + + val rental = RentalCar("Renault", "Ellypse", 2002) + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val violations: Set[ConstraintViolation[RentalCar]] = executableValidator.validateMethods(rental) + assert(violations.size == 1) + assert(violations.head.getMessage == "model year must be within the last 2 years") + assert(violations.head.getPropertyPath.toString == "onlyNewerCars.modelYear") + assert(violations.head.getInvalidValue == rental) + +`ScalaExecutableValidator#validateMethod` ++++++++++++++++++++++++++++++++++++++++++ + +Validates the given `@MethodValidation`-annotated method of an object. + +.. code-block:: scala + :emphasize-lines: 30 + + import com.twitter.util.validation.MethodValidation + import com.twitter.util.validation.engine.MethodValidationResult + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{Min, NotEmpty} + import java.lang.reflect.Method + import java.time.LocalDate + + case class RentalCar( + @NotEmpty make: String, + @NotEmpty model: String, + @Min(2000) modelYear: Int) { + + @MethodValidation(fields = Array("modelYear")) + def onlyNewerCars: MethodValidationResult = { + // only want model years within the last 2 years + val year: Int = LocalDate.now.getYear + if ((year - modelYear) <= 2) MethodValidationResult.Valid() + else MethodValidationResult.Invalid("model year must be within the last 2 years") + } + } + + val rental = RentalCar("Renault", "Nepta", 2006) + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val method = classOf[RentalCar].getDeclaredMethod("onlyNewerCars") + + val violations: Set[ConstraintViolation[RentalCar]] = executableValidator.validateMethod(rental, method) + assert(violations.size == 1) + assert(violations.head.getMessage == "model year must be within the last 2 years") + assert(violations.head.getPropertyPath.toString == "onlyNewerCars.modelYear") + assert(violations.head.getInvalidValue == rental) + +`ScalaExecutableValidator#validateParameters` ++++++++++++++++++++++++++++++++++++++++++++++ + +Validates all constraints placed on the parameters of the given method. This method directly proxies +to the `underlying Hibernate validator `__. + +.. code-block:: scala + :emphasize-lines: 32,33,34,35 + + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{Min, NotEmpty, NotNull, Size} + import java.time.LocalDate + + case class RentalStation(@NotNull name: String) { + def rentCar( + @NotNull customer: Customer, + @NotNull @Future start: LocalDate, + @Min(1) duration: Int + ): Unit = ??? + + @NotEmpty + @Size(min = 1) + def getCustomers: Seq[Customer] = Seq.empty + } + + val rentalStation = RentalStation("Hertz") + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val method = + classOf[RentalStation].getMethod( + "rentCar", + classOf[Customer], + classOf[LocalDate], + classOf[Int]) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array(customer, LocalDate.now(), Integer.valueOf(5))) + assert(violations.size == 1) + assert(violations.head.getMessage == "must be a future date") + assert(violations.head.getPropertyPath.toString == "rentCar.start") + +`ScalaExecutableValidator#validateReturnValue` +++++++++++++++++++++++++++++++++++++++++++++++ + +Validates return value constraints of the given method. This method directly proxies to the +`underlying Hibernate validator `__. + +.. code-block:: scala + :emphasize-lines: 27,28,29,30 + + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{Min, NotEmpty, NotNull, Size} + import java.time.LocalDate + + case class RentalStation(@NotNull name: String) { + def rentCar( + @NotNull customer: Customer, + @NotNull @Future start: LocalDate, + @Min(1) duration: Int + ): Unit = ??? + + @NotEmpty + @Size(min = 1) + def getCustomers: Seq[Customer] = Seq.empty + } + + val rentalStation = RentalStation("Hertz") + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val method = classOf[RentalStation].getMethod("getCustomers") + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateReturnValue( + rentalStation, + method, + Seq.empty[Customer]) + assert(violations.size == 2) + assert(violations.head.getMessage == "must not be empty") + assert(violations.head.getPropertyPath.toString == "getCustomers.") + assert(violations.last.getMessage == "size must be between 1 and 2147483647") + assert(violations.last.getPropertyPath.toString == "getCustomers.") + +`ScalaExecutableValidator#validateConstructorParameters` +++++++++++++++++++++++++++++++++++++++++++++++++++++++++ + +Validates all constraints placed on the parameters of the given constructor. This method directly +proxies to the `underlying Hibernate validator `__. + +.. code-block:: scala + :emphasize-lines: 27,28,29 + + package com.example + + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + import jakarta.validation.constraints.{Min, NotEmpty, NotNull, Size} + import java.time.LocalDate + + case class RentalStation(@NotNull name: String) { + def rentCar( + @NotNull customer: Customer, + @NotNull @Future start: LocalDate, + @Min(1) duration: Int + ): Unit = ??? + + @NotEmpty + @Size(min = 1) + def getCustomers: Seq[Customer] = Seq.empty + } + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateConstructorParameters( + constructor, + Array(null)) + assert(violations.size == 1) + assert(violations.head.getMessage == "must not be null") + assert(violations.head.getPropertyPath.toString == "RentalStation.name") + +`ScalaExecutableValidator#validateConstructorReturnValue` ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ + +Validates a constructor's return value. This method directly proxies to the +`underlying Hibernate validator `__. + +.. code-block:: scala + :emphasize-lines: 25,26,27 + + package com.example + + import com.twitter.util.validation.ScalaValidator + import com.twitter.util.validation.executable.ScalaExecutableValidator + import jakarta.validation.ConstraintViolation + + case class CarWithPassengerCount @ValidPassengerCountReturnValue(max = 1) (passengers: Seq[Person]) + + val car = CarWithPassengerCount( + Seq( + Person("abcd1234", "J Doe", DefaultAddress), + Person("abcd1235", "K Doe", DefaultAddress), + Person("abcd1236", "L Doe", DefaultAddress), + Person("abcd1237", "M Doe", DefaultAddress), + Person("abcd1238", "N Doe", DefaultAddress) + ) + ) + + val validator: ScalaValidator = ScalaValidator() + val executableValidator: ScalaExecutableValidator = validator.forExecutables + + val constructor = classOf[CarWithPassengerCount].getConstructor(classOf[Seq[Person]]) + + val violations: Set[ConstraintViolation[CarWithPassengerCount]] = + executableValidator.validateConstructorReturnValue( + constructor, + car) + assert(violations.size == 1) + assert(violations.head.getMessage == "invalid number of passengers") + assert(violations.head.getPropertyPath.toString == "CarWithPassengerCount.") + +Built-in method constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +In addition to the built-in case class and field-level constraints discussed in +`“Built-in constraints” <#built-in-constraints>`__, this library provides one method-level +constraint, `@MethodValidation`. This is a generic constraint which allows for implementation of +validation of the (possibly even internal) state of a case class instance. + +This is different from a class-level constraint since being implemented as a method within the case +class means that the method implementation potentially has access to private or internal state for +performing validation whereas a class-level constraint only has access to publicly exposed fields +and methods. + +* `@MethodValidation` + + Used to validate the (possibly internal) state of an entire object. It is important to note that any + annotated method **will be invoked when performing validation** and should thus be side-effect free. + + - **Supported data types** + + No-arg case class method that returns a `MethodValidationResult`. + +Grouping constraints +-------------------- + +See the `Hibernate documentation `__ +which explains constraint grouping in detail. + +Creating custom constraints +--------------------------- + +In cases where none of the `built-in constraints <#built-in-constraints>`__ suffice, you can create +custom constraints to implement your specific validation requirements. + +Creating a simple constraint +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +To create a custom constraint, the following three steps are required: + +* Create a constraint annotation +* Implement a `ConstraintValidator `__ +* Define a default error message + +The constraint annotation +^^^^^^^^^^^^^^^^^^^^^^^^^ + +This is detailed in the `Hibernate documentation `__ +and all of the steps here are the same for use with this library. The linked documentation creates a +`CheckCase` annotation with two modes for validating a `String` is either uppercase or lowercase. +We'll assume the `CaseMode` enum and the `CheckCase` constraint annotation exists. + +The constraint validator +^^^^^^^^^^^^^^^^^^^^^^^^ + +Having defined the annotation, we now need to create an implementation of a `ConstraintValidator `__, +which is able to validate fields with a `@CheckCase` annotation. To do so, we implement the Jakarta Bean Validation interface `ConstraintValidator `__: + +.. code-block:: scala + + package com.example.constraint + + import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + + class CheckCaseConstraintValidator extends ConstraintValidator[CheckCase, String] { + @volatile private[this] var caseMode: CaseMode = _ + + override def initialize(constraintAnnotation: CheckCase): Unit = { + this.caseMode = constraintAnnotation.value() + } + + def isValid( + obj: String, + constraintContext: ConstraintValidatorContext): Boolean = { + if (obj == null) true + else { + caseMode match { + case UPPER => + obj == obj.toUpperCase + case LOWER => + obj == obj.toLowerCase + } + } + } + } + +The `ConstraintValidator `__ +interface defines two type parameters which are set in the implementation. The first one specifies +the annotation type to be validated (in this case `CheckCase`), the second one the type of elements +the validator can handle (in this case, `String`). + +In case a constraint supports several data types, you can create a `ConstraintValidator `__ +implementation over the `Any` or `AnyRef` types and pattern-match within the implementation. If the +implementation receives a type that is unsupported you can throw an `UnexpectedTypeException `__. + +The implementation of the validator is straightforward: the `initialize` method gives you access +to the attribute values of the validated constraint and allows you to store them in a field of the +validator as shown in the example. + +The `isValid` method contains the validator logic. For `@CheckCase` this is the check as to whether +a given string is either completely lower case or upper case, depending on the `caseMode` set from +the constraint annotation value in `initialize`. + +.. important:: + + Note that the Jakarta Bean Validation specification recommends consideration of `null` values as + being valid. If `null` is not a valid value for an element, it should be **explicitly** annotated + with `@NotNull`. + + This library also extends the same consideration for `None` values during validation as + detailed in `With Option[_] <#with-option>`__. + +Thread-safety ++++++++++++++ + +As noted in the `javadocs `__: +access to `isValid` can happen concurrently and thus thread-safety must be ensured by the +implementation. Also note that `isValid` implementations **should not** alter the state of the value +to validate. + +The error message +^^^^^^^^^^^^^^^^^ + +The last piece is an error message which should be used when a `@CheckCase` constraint is violated. +You have a few options. You can specify the default error message directly within the `message()` +attribute of the constraint annotation, e.g., + +.. code-block:: java + + String message() default "Case mode must be {value}"; + +However, if you want to be able to `internationalize `__ +an error message, it is generally recommended to place the value in a `Resource Bundle `__, +specifying the resource key as the `message()` attribute in the constraint annotation. + +.. code-block:: java + + String message() default "{com.example.constraint.CheckCase.message}"; + +Create a `properties `__ file, +*ValidationMessages.properties* with the following contents: + +.. code-block:: text + + com.example.constraint.CheckCase.message=Case mode must be {value} + +Make sure the *ValidationMessages.properties* is packaged on your classpath to be picked up by the +library. + +For details, see the Hibernate documentation: `"Default message interpolation" `__. + +The `TwitterConstraintValidatorContext` ++++++++++++++++++++++++++++++++++++++++ + +First, see the documentation about the `ConstraintValidatorContext `__ +for background on its usage within a `ConstraintValidator `__ +implementation. + +This library provides a utility for customizing any returned `ConstraintViolation` from a +`ConstraintValidator` by using the `c.t.util.validation.constraintvalidation.TwitterConstraintValidatorContext`. + +The `TwitterConstraintValidatorContext` uses the `HibernateConstraintValidatorContext `__ +described in the `Custom contexts `__ +Hibernate documentation. See the referenced documentation for details. + +The `TwitterConstraintValidatorContext` has a simple DSL to create custom constraint violations. The +DSL allows you to set a dynamic `Payload `__, +set message parameters or expression variables for an error message, or specify a fully custom error +message template (that is a message template different than the default specified by the constraint +annotation). + +.. note:: + + As mentioned in the Hibernate `documentation `__: + + *[T]he main difference between message parameters and expression variables is that message parameters* + *are simply interpolated whereas expression variables are interpreted using the* `Expression Language engine `__. + + *In practice, use message parameters if you do not need the advanced features of an Expression Language.* + +By default, Expression Language support **is not enabled** for custom violations created via the +`ConstraintValidatorContext`. However, for some advanced requirements, using Expression Language may +be necessary. If you add an expression variable via the `TwitterConstraintValidatorContext`, +Expression Language support will be enabled when adding a constraint violation. + +.. code-block:: scala + :emphasize-lines: 25,26,27,28,29 + + import com.twitter.util.validation.constraintvalidation.TwitterConstraintValidatorContext + import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + + class CheckCaseConstraintValidator extends ConstraintValidator[CheckCase, String] { + @volatile private[this] var caseMode: CaseMode = _ + + override def initialize(constraintAnnotation: CheckCase): Unit = { + this.caseMode = constraintAnnotation.value() + } + + def isValid( + obj: String, + constraintContext: ConstraintValidatorContext): Boolean = { + val valid = if (obj == null) true + else { + caseMode match { + case UPPER => + obj == obj.toUpperCase + case LOWER => + obj == obj.toLowerCase + } + } + + if (!valid) { + TwitterConstraintValidatorContext + .withDynamicPayload(???) + .addMessageParameter("foo", "snafu") + .withMessageTemplate("{foo}") // will be interpolated into the string 'snafu' + .addConstraintViolation(constraintContext) + } + + valid + } + } + +Using the constraint +^^^^^^^^^^^^^^^^^^^^ + +To use the custom `@CheckCase` constraint, annotate a case class field: + +.. code-block:: scala + :emphasize-lines: 7 + + package com.example + + import jakarta.validation.constraints.{Min, NotEmpty, Size} + + case class Car( + @NotEmpty manufacturer: String, + @NotEmpty @Size(min = 2, max = 14) @CheckCase(CaseMode.UPPER) licensePlate: String, + @Min(2) seatCount: Int) + +During validation, a `licensePlate` field with an invalid case will now cause validation to fail. + +.. code-block:: scala + + // invalid license plate + val car: Car = Car("Morris", "dd-ab-123", 4) + + val violations: Set[ConstraintViolation[Car]] = validator.validate(car) + assert(violations.size == 1) + assert(violations.head.getMessage == "Case mode must be UPPER") + + // valid license plate + val car: Car = Car("Morris", "DD-AB-123", 4) + assert(validator.validate(car).size == 0) + +Class-level constraints +~~~~~~~~~~~~~~~~~~~~~~~ + +As mentioned previously, constraints can also be applied on the class-level to validate the state of +the entire case class instance. These constraints are defined in the same way as field constraints. + +Here we'll see how to implement the `@ValidPassengerCount` annotation seen in an earlier example: + +.. code-block:: scala + :emphasize-lines: 6 + + package com.example + + import com.example.constraint.ValidPassengerCount + import jakarta.validation.constraints.Min + + @ValidPassengerCount + case class Car( + @Min(2) seatCount: Int, + passengers: Seq[Person]) + +First the annotation definition could be written as: + +.. code-block:: java + + package com.example.constraint; + + @Target({ TYPE, ANNOTATION_TYPE }) + @Retention(RUNTIME) + @Constraint(validatedBy = { ValidPassengerCountValidator.class }) + @Documented + public @interface ValidPassengerCount { + + String message() default "{com.example.constraint.ValidPassengerCount.message}"; + + Class[] groups() default { }; + + Class[] payload() default { }; + } + +We define an error message in a *Validation.properties* file: + +.. code-block:: text + + com.example.constraint.ValidPassengerCount.message=invalid number of passengers + +Then the `ConstraintValidator` implementation which is specified to a `Car` type: + +.. code-block:: scala + + package com.example + + import com.example.constraint.ValidPassengerCount + import jakarta.validation.{ConstraintValidator + + class ValidPassengerCountValidator extends ConstraintValidator[ValidPassengerCount, Car] { + + def isValid( + car: Car, + constraintContext: ConstraintValidatorContext + ): Boolean = + if (car == null) true + else car.passengers.size <= car.seatCount + } + +Note as the example shows, you need to use the element type `TYPE` in the `@Target `__ +annotation on the constraint annotation. This allows the constraint to be put on type definitions. +The `ValidPassengerCountValidator` receives an instance of a `Car` in the `isValid()` method and can +access the complete object state to decide whether the given instance is valid or not. + +During validation, an instance which does not satisfy the constraint will fail the validation. + +.. code-block:: scala + + package com.example + + import com.twitter.util.validation.ScalaValidator + import jakarta.validation.ConstraintViolation + + val car = Car( + seatCount = 2, + passengers = Seq( + Person("abcd1234", "J Doe", DefaultAddress), + Person("abcd1235", "K Doe", DefaultAddress), + Person("abcd1236", "L Doe", DefaultAddress), + Person("abcd1237", "M Doe", DefaultAddress), + Person("abcd1238", "N Doe", DefaultAddress) + ) + ) + + val validator: ScalaValidator = ScalaValidator() + + val violations: Set[ConstraintViolation[Car]] = + validator.validate(car) + assert(violations.size == 1) + assert(violations.head.getMessage == "invalid number of passengers") + assert(violations.head.getPropertyPath.toString == "Car") + +Custom property paths +^^^^^^^^^^^^^^^^^^^^^ + +By default for a class-level constraint, any constraint violation is reported on the level of the +annotated type, e.g. in this case the `Car` type. If the class is the top-level, the property path +will be empty, otherwise it will be the field name referencing the class type. + +See the `Hibernate documentation `__ +for details on how to set custom property paths when implementing a class-level constraint. + +Cross-parameter constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The Bean validation specification distinguishes two types of constraints: generic constraints which +apply to the annotated element (e.g., a type, field, container element, method parameter or return +value, etc.), and *cross-parameter* constraints. Cross-parameter constraints by contrast, apply to +an array of parameters of a method or constructor and can be used to express validation logic which +depends on several parameter values. + +See the `Hibernate documentation `__ +for details on implementing cross-parameter constraints. + +Constraint composition +~~~~~~~~~~~~~~~~~~~~~~ + +To keep the annotation of case class fields from becoming unwieldy or confusing, you can use a +composing constraint. See the `Hibernate documentation `__ +for details on creating a composing constraint. + +`@MethodValidation` constraints +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +.. important:: + + The `@MethodValidation` constraint annotation MUST be applied to a no-arg method that returns a + `com.twitter.util.validation.engine.MethodValidationResult`. + + These annotated methods **will be invoked** during validation and should be side-effect free. + +In a case class, define a no-arg method that returns a `MethodValidationResult`. Annotate the method +with `@MethodValidation`, optionally specifying validated fields of the case class for +error-reporting. + +.. code-block:: scala + :emphasize-lines: 11,12 + + import com.twitter.util.validation.MethodValidation + import com.twitter.util.validation.engine.MethodValidationResult + import jakarta.validation.constraints.{Min, NotEmpty} + import java.time.LocalDate + + case class RentalCar( + @NotEmpty make: String, + @NotEmpty model: String, + @Min(2000) modelYear: Int) { + + @MethodValidation(fields = Array("modelYear")) + def onlyNewerCars: MethodValidationResult = { + // only want model years within the last 2 years + val year: Int = LocalDate.now.getYear + if ((year - modelYear) <= 2) MethodValidationResult.Valid() + else MethodValidationResult.Invalid("model year must be within the last 2 years") + } + } + +Value Extraction +---------------- + +Value extraction is the process of extracting values from a container so that they can be validated. + +It is used when dealing with `container element constraints <#container-element-constraints>`__ and +`cascaded validation <#object-graphs>`__ inside containers. + +Built-in value extractors +~~~~~~~~~~~~~~~~~~~~~~~~~ + +This library provides built-in value extraction for Scala `Iterable[_] <#iterable>`__ types. See the +`Hibernate documentation `__ +for more details on value extraction including how to define and register a new value extractor. + +Best Practices +-------------- + +- Case classes used for validation should be side-effect free under property access. That is, accessing a + case class field should be able to be done eagerly and without side-effects. +- Case class methods for `@MethodValidation` validation should be side-effect free under method access. That is, accessing a + `@MethodValidation`-annotated method should be able to be done eagerly and without side-effects. +- Case classes with generic type params should be considered **not supported** for validation. E.g., + + .. code-block:: scala + + case class GenericCaseClass[T](@NotEmpty @Valid data: T) + + This may appear to work for some category of data but in practice the way the library caches reflection data does not + discriminate on the type binding and thus validation of genericized case classes is not fully guaranteed to be successful for + differing type bindings of a given genericized class. +- While `Iterable` collections are supported for cascaded validation when annotated with `@Valid`, this does + not include `Map` types. Annotated `Map` types will be ignored during validation. +- More generally, types with multiple type params, e.g. `Either[T, U]`, are not supported for validation + of contents when annotated with `@Valid`. Annotated unsupported types will be **ignored** during validation. + +More resources +-------------- + +* `Hibernate Validator Specifics `__ (from reference guide) +* `Hibernate Validator Reference Guide `__ +* `Jakarta Bean Validation API 3.0.0 Javadocs `__ +* `Hibernate Validator API Javadocs `__ +* `Jakarta Bean Validation `__ diff --git a/project/build.properties b/project/build.properties index 23050492c6..22af2628c4 100644 --- a/project/build.properties +++ b/project/build.properties @@ -1 +1 @@ -sbt.version=1.1.4 \ No newline at end of file +sbt.version=1.7.1 diff --git a/project/build.sbt b/project/build.sbt index 0c8c057a88..26f5fb29cc 100644 --- a/project/build.sbt +++ b/project/build.sbt @@ -2,5 +2,6 @@ scalacOptions ++= Seq( "-deprecation", "-unchecked", "-feature", - "-encoding", "utf8" + "-encoding", + "utf8" ) diff --git a/project/plugins.sbt b/project/plugins.sbt index 51c141d1b1..c6af984fd4 100644 --- a/project/plugins.sbt +++ b/project/plugins.sbt @@ -1,7 +1,6 @@ resolvers += Classpaths.sbtPluginReleases resolvers += Resolver.sonatypeRepo("snapshots") -addSbtPlugin("com.typesafe.sbt" % "sbt-site" % "1.3.0") -addSbtPlugin("io.get-coursier" % "sbt-coursier" % "1.0.0-RC13") -addSbtPlugin("org.scoverage" % "sbt-scoverage" % "1.5.1") -addSbtPlugin("pl.project13.scala" % "sbt-jmh" % "0.3.4") +addSbtPlugin("com.typesafe.sbt" % "sbt-site" % "1.4.1") +addSbtPlugin("pl.project13.scala" % "sbt-jmh" % "0.4.3") +addSbtPlugin("ch.epfl.scala" % "sbt-scalafix" % "0.10.1") diff --git a/project/unidoc.sbt b/project/unidoc.sbt index 046b1130a6..56bff7630e 100644 --- a/project/unidoc.sbt +++ b/project/unidoc.sbt @@ -1,2 +1 @@ -addSbtPlugin("com.eed3si9n" % "sbt-unidoc" % "0.4.1") - +addSbtPlugin("com.github.sbt" % "sbt-unidoc" % "0.5.0") diff --git a/pushsite.bash b/pushsite.bash index dab72393d9..a3d6a8ceb0 100755 --- a/pushsite.bash +++ b/pushsite.bash @@ -1,4 +1,4 @@ -#!/bin/bash +#!/usr/bin/env bash set -e diff --git a/sbt b/sbt index 9e21a00f08..dbe5a02103 100755 --- a/sbt +++ b/sbt @@ -1,21 +1,23 @@ -#!/bin/bash +#!/usr/bin/env bash -sbtver=1.1.4 -sbtjar=sbt-launch.jar -sbtsha128=ac6b13b709e3c0e3a8b2b3bf3031ec57660cb813 +set -eo pipefail + +sbtver="1.7.1" +sbtjar="sbt-launch-$sbtver.jar" +sbtsha128="468efdd45baf58dbb575f9b4369c5234f8cd54ba" sbtrepo="https://repo1.maven.org/maven2/org/scala-sbt/sbt-launch" -if [ ! -f $sbtjar ]; then +if [ ! -f "sbt-launch.jar" ]; then echo "downloading $PWD/$sbtjar" 1>&2 - if ! curl --location --silent --fail --remote-name $sbtrepo/$sbtver/$sbtjar; then + if ! curl --location --silent --fail -o "sbt-launch.jar" "$sbtrepo/$sbtver/$sbtjar"; then exit 1 fi fi -checksum=`openssl dgst -sha1 $sbtjar | awk '{ print $2 }'` +checksum=`openssl dgst -sha1 sbt-launch.jar | awk '{ print $2 }'` if [ "$checksum" != $sbtsha128 ]; then - echo "bad $PWD/$sbtjar. delete $PWD/$sbtjar and run $0 again." + echo "bad $PWD/sbt-launch.jar. delete $PWD/sbt-launch.jar and run $0 again." exit 1 fi @@ -24,17 +26,4 @@ fi java -ea \ $SBT_OPTS \ $JAVA_OPTS \ - -Djava.net.preferIPv4Stack=true \ - -XX:+AggressiveOpts \ - -XX:+UseParNewGC \ - -XX:+UseConcMarkSweepGC \ - -XX:+CMSParallelRemarkEnabled \ - -XX:+CMSClassUnloadingEnabled \ - -XX:ReservedCodeCacheSize=128m \ - -XX:SurvivorRatio=128 \ - -XX:MaxTenuringThreshold=0 \ - -Xss8M \ - -Xms512M \ - -Xmx2G \ - -server \ - -jar $sbtjar "$@" + -jar "sbt-launch.jar" "$@" diff --git a/site/index.html b/site/index.html index cd4e154e43..bc308a8316 100644 --- a/site/index.html +++ b/site/index.html @@ -69,7 +69,7 @@

Util

  • User’s guides
  • API documentation
  • Finaglers Gitter channel
  • -
  • Finaglers Google Group
  • +
  • Finaglers Google Group
  • Contributing

    @@ -85,7 +85,7 @@

    Contributing

    Before endeavoring on large changes, please discuss them with the - Finagle Google Group + Finagle Google Group to receive feedback and suggestions.

    For all patches, please review our diff --git a/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/BUILD b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/BUILD new file mode 100644 index 0000000000..5263e66e3c --- /dev/null +++ b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/BUILD @@ -0,0 +1,18 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-app-lifecycle", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], + exports = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Event.scala b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Event.scala new file mode 100644 index 0000000000..bf10e145fa --- /dev/null +++ b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Event.scala @@ -0,0 +1,46 @@ +package com.twitter.app.lifecycle + +/** + * Represents the phase in the application lifecycle that an event occurs in. + * + * App's Lifecycle comprises the following phases + * + * [ Register ] -> [ Load Bindings ] -> [ Init ] -> [ Parse Args ] -> [ Pre Main ] -> + * [ Main ] -> [ Post Main ] -> [ Close ] -> [ CloseExitLast ] + * + * @see [[https://twitter.github.io/twitter-server/Features.html#lifecycle-management TwitterServer]] + * @see [[https://twitter.github.io/finatra/user-guide/getting-started/lifecycle.html Finatra]] + */ +private[twitter] sealed trait Event + +private[twitter] object Event { + + /* com.twitter.app.App + com.twitter.server.Server events */ + case object Register extends Event + case object LoadBindings extends Event + case object Init extends Event + case object ParseArgs extends Event + case object PreMain extends Event + case object Main extends Event + case object PostMain extends Event + case object Close extends Event + case object CloseExitLast extends Event + + /* com.twitter.server.Lifecycle.Warmup events */ + case object PrebindWarmup extends Event + case object WarmupComplete extends Event + + /* com.twitter.inject.App + com.twitter.inject.TwitterServer events */ + /* + * [ Register ] -> [ Load Bindings] -> [ Init ] -> [ Parse Args ] -> [ Pre Main ] -> + * [ Main // the following events take place within main() + * ( Post Injector Startup ) -> ( Warmup ) -> ( Before Post Warmup ) -> ( Post Warmup ) -> + * ( After Post Warmup ) + * ] -> [ Post Main ] -> [ Close ] -> [ CloseExitLast ] + */ + case object PostInjectorStartup extends Event + case object Warmup extends Event + case object BeforePostWarmup extends Event + case object PostWarmup extends Event + case object AfterPostWarmup extends Event +} diff --git a/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Notifier.scala b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Notifier.scala new file mode 100644 index 0000000000..5424961aa4 --- /dev/null +++ b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Notifier.scala @@ -0,0 +1,64 @@ +package com.twitter.app.lifecycle + +import com.twitter.util.{Future, Return, Throw} + +/** + * The [[Notifier]] emits signals to [[Observer observers]] when lifecycle [[Event events]] are + * executed. + * + * @param observers the observers who will be notified of lifecycle event activities + * @note exposed for testing + */ +private[twitter] class Notifier private[app] (observers: Iterable[Observer]) { + + /** + * Record and execute a lifecycle event and emit the result to any [[Observer observers]]. + * + * @param event The [[Event lifecycle event]] to execute. + * @param f The function to execute as part of this [[Event lifecycle event]]. + */ + final def apply(event: Event)(f: => Unit): Unit = execAndSignal(event, f) + + /** + * Record and execute a lifecycle event [[Future]] and emit the result to + * any [[Observer observers]] on completion of the [[Future]]. + * + * @param event The [[Event lifecycle event]] to execute. + * @param f The Future that will be responded to on completion for this + * [[Event lifecycle event]]. + */ + final def future(event: Event)(f: Future[Unit]): Future[Unit] = { + onEntry(event) + f.respond { + case Return(_) => + onSuccess(event) + case Throw(t) => + onFailure(event, t) + } + } + + private[this] def execAndSignal(event: Event, f: => Unit): Unit = { + onEntry(event) + try { + f + onSuccess(event) + } catch { + case t: Throwable => + onFailure(event, t) + throw t + } + } + + // notify all observers we are entering the event + private[this] def onEntry(stage: Event): Unit = + observers.foreach(_.onEntry(stage)) + + // notify all observers on success + private[this] def onSuccess(stage: Event): Unit = + observers.foreach(_.onSuccess(stage)) + + // notify all observers on failures + private[this] def onFailure(stage: Event, throwable: Throwable): Unit = + observers.foreach(_.onFailure(stage, throwable)) + +} diff --git a/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Observer.scala b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Observer.scala new file mode 100644 index 0000000000..f29788a812 --- /dev/null +++ b/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle/Observer.scala @@ -0,0 +1,16 @@ +package com.twitter.app.lifecycle + +/** + * A trait for observing lifecycle event notifications + */ +private[twitter] trait Observer { + + /** notification that the [[event]] lifecycle phase has begun */ + def onEntry(event: Event): Unit = () + + /** notification that the [[event]] lifecycle phase has completed successfully */ + def onSuccess(event: Event): Unit + + /** notification that the [[event]] lifecycle phase has encountered [[throwable]] error */ + def onFailure(stage: Event, throwable: Throwable): Unit +} diff --git a/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/BUILD.bazel b/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/BUILD.bazel new file mode 100644 index 0000000000..ec4032ff29 --- /dev/null +++ b/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/BUILD.bazel @@ -0,0 +1,17 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle", + "util/util-core:util-core-util", + ], +) diff --git a/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/NotifierTest.scala b/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/NotifierTest.scala new file mode 100644 index 0000000000..4b07284520 --- /dev/null +++ b/util-app-lifecycle/src/test/scala/com/twitter/app/lifecycle/NotifierTest.scala @@ -0,0 +1,94 @@ +package com.twitter.app.lifecycle + +import Event._ +import org.scalatest.funsuite.AnyFunSuite + +class NotifierTest extends AnyFunSuite { + + test("emits success") { + var success = false + var observedEvent: Event = null + val observer = new Observer { + def onSuccess(event: Event): Unit = { + observedEvent = event + success = true + } + + def onFailure(event: Event, throwable: Throwable): Unit = { + observedEvent = event + success = false + } + } + val notifier = new Notifier(Seq(observer)) + notifier(PreMain) { + // success + () + } + assert(observedEvent == PreMain, "Observer did not see PreMain event") + assert(success, "Observer did not see onSuccess") + } + + test("emits failure and throws") { + var failed = false + var observedEvent: Event = null + var thrown: Throwable = null + + val observer = new Observer { + def onSuccess(event: Event): Unit = { + observedEvent = event + failed = false + } + + def onFailure(event: Event, throwable: Throwable): Unit = { + observedEvent = event + failed = true + thrown = throwable + } + } + val notifier = new Notifier(Seq(observer)) + + intercept[IllegalStateException] { + notifier(PostMain) { + throw new IllegalStateException("BOOM!") + } + } + } + + test("continues to emit after failure") { + var failed = false + var observedEvent: Event = null + var thrown: Throwable = null + + val observer = new Observer { + def onSuccess(event: Event): Unit = { + observedEvent = event + failed = false + } + + def onFailure(event: Event, throwable: Throwable): Unit = { + observedEvent = event + failed = true + thrown = throwable + } + } + val notifier = new Notifier(Seq(observer)) + + intercept[IllegalStateException] { + notifier(PostMain) { + throw new IllegalStateException("BOOM!") + } + } + assert(observedEvent == PostMain, "Observer did not see PostMain event") + assert(failed, "Observer did not see onFailure") + assert(thrown.isInstanceOf[IllegalStateException] && thrown.getMessage == "BOOM!") + + // our next lifecycle event will succeed + notifier(Close) { + () + } + + assert(observedEvent == Close) + assert(failed == false) + } + +} diff --git a/util-app/BUILD b/util-app/BUILD index e26545d608..396a12cf8f 100644 --- a/util-app/BUILD +++ b/util-app/BUILD @@ -1,12 +1,14 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-app/src/main/java", "util/util-app/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-app/src/test/java", "util/util-app/src/test/scala", diff --git a/util-app/OWNERS b/util-app/OWNERS deleted file mode 100644 index 463d61c17c..0000000000 --- a/util-app/OWNERS +++ /dev/null @@ -1,6 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta -ryano diff --git a/util-app/PROJECT b/util-app/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-app/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-app/src/main/java/BUILD b/util-app/src/main/java/BUILD index 5441291046..dcbe740001 100644 --- a/util-app/src/main/java/BUILD +++ b/util-app/src/main/java/BUILD @@ -1,11 +1,7 @@ -java_library( - sources = rglobs("*.java"), - fatal_warnings = True, - provides = artifact( - org = "com.twitter", - name = "util-app-java", - repo = artifactory, - ), - # Not detected as required by zinc. - scope = "forced", +target( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-app/src/main/java/com/twitter/app", + "util/util-app/src/main/scala", + ], ) diff --git a/util-app/src/main/java/com/twitter/app/BUILD b/util-app/src/main/java/com/twitter/app/BUILD new file mode 100644 index 0000000000..14f71e72de --- /dev/null +++ b/util-app/src/main/java/com/twitter/app/BUILD @@ -0,0 +1,13 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-app-java", + repo = artifactory, + ), + # Not detected as required by zinc. + scope = "forced", + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-app/src/main/scala/BUILD b/util-app/src/main/scala/BUILD index 4a7dd01433..8dd879de25 100644 --- a/util-app/src/main/scala/BUILD +++ b/util-app/src/main/scala/BUILD @@ -1,18 +1,28 @@ scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", provides = scala_artifact( org = "com.twitter", name = "util-app", repo = artifactory, ), + strict_deps = True, + tags = ["bazel-compatible"], dependencies = [ - "util/util-app/src/main/java", - "util/util-core/src/main/scala", - "util/util-registry/src/main/scala", + "util/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle", + "util/util-app/src/main/java/com/twitter/app", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-registry/src/main/scala/com/twitter/util/registry", ], exports = [ - "util/util-app/src/main/java", - "util/util-core/src/main/scala", + "util/util-app-lifecycle/src/main/scala/com/twitter/app/lifecycle", + "util/util-app/src/main/java/com/twitter/app", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", ], ) diff --git a/util-app/src/main/scala/com/twitter/app/App.scala b/util-app/src/main/scala/com/twitter/app/App.scala index fbd03ba8e2..27f6fd50e0 100644 --- a/util-app/src/main/scala/com/twitter/app/App.scala +++ b/util-app/src/main/scala/com/twitter/app/App.scala @@ -1,14 +1,18 @@ package com.twitter.app -import com.twitter.conversions.time._ +import com.twitter.app.lifecycle.Event._ +import com.twitter.conversions.DurationOps._ import com.twitter.util._ +import java.io.PrintWriter +import java.io.StringWriter import java.lang.reflect.InvocationTargetException import java.util.concurrent.ConcurrentLinkedQueue import java.util.concurrent.atomic.AtomicReference -import java.util.logging.Logger -import scala.collection.JavaConverters._ import scala.collection.mutable +import scala.jdk.CollectionConverters._ +import scala.sys.ShutdownHookThread import scala.util.control.NonFatal +import com.twitter.finagle.util.enableJvmShutdownHook /** * A composable application trait that includes flag parsing as well @@ -35,16 +39,22 @@ import scala.util.control.NonFatal * * Note that a missing `main` is OK: mixins may provide behavior that * does not require defining a custom `main` method. + * + * @define LoadServiceApplyScaladocLink + * [[com.twitter.app.LoadService.apply[T]()(implicitevidence\$1:scala\.reflect\.ClassTag[T]):Seq[T]* LoadService.apply]] */ -trait App extends Closable with CloseAwaitably { +trait App extends ClosableOnce with CloseOnceAwaitably with Lifecycle { /** The name of the application, based on the classname */ - val name: String = getClass.getName.stripSuffix("$") + def name: String = getClass.getName.stripSuffix("$") - /** The [[com.twitter.app.Flags]] instance associated with this application */ //failfastOnFlagsNotParsed is called in the ctor of App.scala here which is a bad idea - //as things like this can happen http://stackoverflow.com/questions/18138397/calling-method-from-constructor - val flag: Flags = new Flags(name, includeGlobal = true, failfastOnFlagsNotParsed) + //as things like this can happen https://stackoverflow.com/questions/18138397/calling-method-from-constructor + private[this] val _flag: Flags = + new Flags(name, includeGlobal = includeGlobalFlags, failfastOnFlagsNotParsed) + + /** The [[com.twitter.app.Flags]] instance associated with this application */ + final def flag: Flags = _flag private var _args = Array[String]() @@ -54,6 +64,9 @@ trait App extends Closable with CloseAwaitably { /** Whether or not to accept undefined flags */ protected def allowUndefinedFlags: Boolean = false + /** Whether or not to include [[GlobalFlag GlobalFlag's]] when [[Flags]] are parsed */ + protected def includeGlobalFlags: Boolean = true + /** * Users of this code should override this to `true` so that * you fail-fast instead of being surprised at runtime by code that @@ -66,14 +79,20 @@ trait App extends Closable with CloseAwaitably { /** Exit on error with the given Throwable */ protected def exitOnError(throwable: Throwable): Unit = { - throwable.printStackTrace() + throwable.printStackTrace(System.out) throwable match { case _: CloseException => // exception occurred while closing, do not attempt to close again System.err.println(throwable.getMessage) System.exit(1) case _ => - exitOnError("Exception thrown in main on startup") + util.Using(new StringWriter) { stringWriter => + util.Using(new PrintWriter(stringWriter)) { printWriter => + throwable.printStackTrace(printWriter) + printWriter.flush() + exitOnError(s"Exception thrown in main on startup: ${stringWriter.toString}") + } + } } } @@ -93,7 +112,13 @@ trait App extends Closable with CloseAwaitably { // want to ensure "reason" is written before attempting to write any details System.err.flush() System.err.println(details) - Await.ready(close(), closeDeadline - Time.now) + // if we have gotten here we may not yet have attempted to close, ensure we do. + try { + Await.ready(close(), (closeDeadline + SecondChanceGrace).sinceNow) + } catch { + case NonFatal(exc) => + exitOnError(newCloseException(Seq(exc))) + } System.exit(1) } @@ -115,7 +140,8 @@ trait App extends Closable with CloseAwaitably { protected lazy val shutdownTimer: Timer = new JavaTimer(isDaemon = true) /** - * Programmatically specify which implementations to use in [[LoadService.apply]] + * Programmatically specify which implementations to use in + * $LoadServiceApplyScaladocLink * for the given interfaces. This allows applications to circumvent * the standard service loading mechanism when needed. It may be useful * if the application has a broad and/or rapidly changing set of dependencies. @@ -142,7 +168,7 @@ trait App extends Closable with CloseAwaitably { * } * }}} * - * If this is called for a `Class[T]` after [[LoadService.apply]] + * If this is called for a `Class[T]` after $LoadServiceApplyScaladocLink * has been called for that same interface, an `IllegalStateException` * will be thrown. For this reason, bindings are done as early * as possible in the application lifecycle, before even `inits` @@ -152,9 +178,9 @@ trait App extends Closable with CloseAwaitably { * their user's implementation choice. * * @return a mapping from `Class` to the 1 or more implementations to - * be used by [[LoadService.apply]] for that interface. + * be used by $LoadServiceApplyScaladocLink for that interface. * - * @see [[LoadService.apply]] + * @see $LoadServiceApplyScaladocLink */ protected[this] def loadServiceBindings: Seq[LoadService.Binding[_]] = Nil @@ -195,6 +221,13 @@ trait App extends Closable with CloseAwaitably { /** Minimum duration to allow for exits to be processed. */ final val MinGrace: Duration = 1.second + /** + * Each resource is given a "grace period" to close gracefully. If it's not able to close within + * the allowed time, we give a "second chance" and wait a little more before forcefully + * terminating the closing. + */ + private[this] final val SecondChanceGrace: Duration = 1.second + /** * Default amount of time to wait for shutdown. * This value is not used as a default if `close()` is called without parameters. It simply @@ -207,23 +240,16 @@ trait App extends Closable with CloseAwaitably { */ @volatile private[this] var closeDeadline = Time.Top - /** - * This is satisfied when all members of `exits` and `lastExits` have closed. - * - * @note Access needs to be mediated via the intrinsic lock. - */ - private[this] var closing: Future[Unit] = Future.never - /** * Close `closable` when shutdown is requested. Closables are closed in parallel. */ final def closeOnExit(closable: Closable): Unit = synchronized { - if (closing == Future.never) { + if (isClosed) { + // `close()` has already been called, we close eagerly + closable.close(closeDeadline) + } else { // `close()` not yet called, safe to add it exits.add(closable) - } else { - // `close()` already called, we need to close this here, immediately. - closable.close(closeDeadline) } } @@ -239,17 +265,15 @@ trait App extends Closable with CloseAwaitably { * In all cases, the close deadline is enforced. */ final def closeOnExitLast(closable: Closable): Unit = synchronized { - if (closing == Future.never) { - // `close()` not yet called, safe to add it - lastExits.add(closable) - } else { + if (isClosed) { // `close()` already called, we need to close this here, but only // after `close()` completes and `closing` is satisfied - closing - .transform { _ => - closable.close(closeDeadline) - } - .by(shutdownTimer, closeDeadline) + close(closeDeadline) + .transform { _ => closable.close(closeDeadline) } + .by(shutdownTimer, closeDeadline + SecondChanceGrace) + } else { + // `close()` not yet called, safe to add it + lastExits.add(closable) } } @@ -257,6 +281,8 @@ trait App extends Closable with CloseAwaitably { * Invoke `f` when shutdown is requested. Exit hooks run in parallel and are * executed after all postmains complete. The thread resumes when all exit * hooks complete or `closeDeadline` expires. + * + * @see [[runOnExit(Runnable)]] for a variant that is more suitable for Java. */ protected final def onExit(f: => Unit): Unit = { closeOnExit { @@ -266,6 +292,67 @@ trait App extends Closable with CloseAwaitably { } } + /** + * Invoke `runnable.run()` when shutdown is requested. Exit hooks run in parallel + * and are executed after all postmains complete. The thread resumes when all exit + * hooks complete or `closeDeadline` expires. + * + * This is a Java friendly API to allow Java 8 users to invoke `onExit` with lambda + * expressions. + * + * For example (in Java): + * + * {{{ + * runOnExit(() -> { + * clientA.close(); + * clientB.close(); + * }); + * }}} + * + * @see [[onExit(f: => Unit)]] for a variant that is more suitable for Scala. + */ + protected final def runOnExit(runnable: Runnable): Unit = { + onExit(runnable.run()) + } + + /** + * Invoke `f` when shutdown is requested. Exit hooks run in parallel and are + * executed after all closeOnExit functions complete. The thread resumes when all exit + * hooks complete or `closeDeadline` expires. + * + * @see [[runOnExitLast(Runnable)]] for a variant that is more suitable for Java. + */ + protected final def onExitLast(f: => Unit): Unit = { + closeOnExitLast { + Closable.make { deadline => // close() ensures that this deadline is sane + FuturePool.unboundedPool(f).by(shutdownTimer, deadline) + } + } + } + + /** + * Invoke `runnable.run()` when shutdown is requested. Exit hooks run in parallel + * and are executed after all closeOnExit functions complete. The thread resumes + * when all exit hooks complete or `closeDeadline` expires. + * + * This is a Java friendly API to allow Java 8 users to invoke `onExitLast` with lambda + * expressions. + * + * For example (in Java): + * + * {{{ + * runOnExitLast(() -> { + * clientA.close(); + * clientB.close(); + * }); + * }}} + * + * @see [[onExitLast(f: => Unit)]] for a variant that is more suitable for Scala. + */ + protected final def runOnExitLast(runnable: Runnable): Unit = { + onExitLast(runnable.run()) + } + /** * Invoke `f` after the user's main has exited. */ @@ -273,26 +360,23 @@ trait App extends Closable with CloseAwaitably { postmains.add(() => f) } - /** - * Notify the application that it may stop running. - * Returns a Future that is satisfied when the App has been torn down or errors at the deadline. - */ - final def close(deadline: Time): Future[Unit] = synchronized { - closing = closeAwaitably { - closeDeadline = deadline.max(Time.now + MinGrace) - Future - .collectToTry(exits.asScala.toSeq.map(_.close(closeDeadline))) - .by(shutdownTimer, closeDeadline) - .transform { - case Return(results) => + override protected def closeOnce(deadline: Time): Future[Unit] = synchronized { + closeDeadline = deadline.max(Time.now + MinGrace) + Future + .collectToTry(exits.asScala.toSeq.map(_.close(closeDeadline))) + .by(shutdownTimer, closeDeadline + SecondChanceGrace) + .transform { + case Return(results) => + observeFuture(CloseExitLast) { closeLastExits(results, closeDeadline) - case Throw(t) => - // this would be triggered by a timeout on the collectToTry of exits, - // still try to close last exits + } + case Throw(t) => + // this would be triggered by a timeout on the collectToTry of exits, + // still try to close last exits + observeFuture(CloseExitLast) { closeLastExits(Seq(Throw(t)), closeDeadline) - } - } - closing + } + } } private[this] def newCloseException(errors: Seq[Throwable]): CloseException = { @@ -302,9 +386,7 @@ trait App extends Closable with CloseAwaitably { s"${errors.size} errors occurred on exit" } val exc = new CloseException(message) - errors.foreach { error => - exc.addSuppressed(error) - } + errors.foreach { error => exc.addSuppressed(error) } exc } @@ -314,7 +396,7 @@ trait App extends Closable with CloseAwaitably { ): Future[Unit] = { Future .collectToTry(lastExits.asScala.toSeq.map(_.close(deadline))) - .by(shutdownTimer, deadline) + .by(shutdownTimer, deadline + SecondChanceGrace) .transform { case Return(results) => val errors = (onExitResults ++ results).collect { case Throw(e) => e } @@ -342,50 +424,71 @@ trait App extends Closable with CloseAwaitably { } final def nonExitingMain(args: Array[String]): Unit = { - App.register(this) - - loadServiceBindings.foreach { binding => - LoadService.bind(binding) + observe(Register) { + App.register(this) } - - for (f <- inits) f() - - parseArgs(args) - - for (f <- premains) f() - - // Get a main() if it's defined. It's possible to define traits that only use pre/post mains. - val mainMethod = - try Some(getClass.getMethod("main")) - catch { case _: NoSuchMethodException => None } - - // Invoke main() if it exists. - mainMethod.foreach { method => - try method.invoke(this) - catch { case e: InvocationTargetException => throw e.getCause } + observe(LoadBindings) { + loadServiceBindings.foreach { binding => LoadService.bind(binding) } } + observe(Init) { + for (f <- inits) f() + } + observe(ParseArgs) { + parseArgs(args) + } + observe(PreMain) { + for (f <- premains) f() + } + observe(Main) { + // Get a main() if it's defined. It's possible to define traits that only use pre/post mains. + val mainMethod = + try Some(getClass.getMethod("main")) + catch { + case _: NoSuchMethodException => None + } + if (enableJvmShutdownHook()) { + ShutdownHookThread { + Await.result(close(defaultCloseGracePeriod)) + } + } - for (f <- postmains.asScala) f() - - // We discard this future but we `Await.result` on `this` which is a - // `CloseAwaitable`, and this means the thread waits for the future to - // complete. - close(defaultCloseGracePeriod) + // Invoke main() if it exists. + mainMethod.foreach { method => + try method.invoke(this) + catch { + case e: InvocationTargetException => throw e.getCause + } + } + } + observe(PostMain) { + for (f <- postmains.asScala) f() + } - // The deadline to 'close' is advisory; we enforce it here. - if (!suppressGracefulShutdownErrors) Await.result(this, closeDeadline - Time.now) - else { - try { // even if we suppress shutdown errors, we give the resources time to close - Await.ready(this, closeDeadline - Time.now) - } catch { - case NonFatal(_) => () + observe(Close) { + // We get a reference to the `close()` Future in order to Await upon it, to ensure the thread + // waits for `close()` to complete. + // Note that the Future returned here is the same as `promise` exposed via ClosableOnce, + // which is the same result that would be returned by an `Await.result` on `this`. + val closeFuture = close(defaultCloseGracePeriod) + + val waitForIt = (closeDeadline + SecondChanceGrace).sinceNow + + // No need to pass a timeout to Await, close(...) enforces the timeout for us. + if (!suppressGracefulShutdownErrors) Await.result(closeFuture, waitForIt) + else { + try { // even if we suppress shutdown errors, we give the resources time to close + Await.ready(closeFuture, waitForIt) + } catch { + case e: TimeoutException => + throw e // we want TimeoutExceptions to propagate + case NonFatal(_) => () + } } } } } object App { - private[this] val log = Logger.getLogger(getClass.getName) private[this] val ref = new AtomicReference[Option[App]](None) /** @@ -395,10 +498,5 @@ object App { */ def registered: Option[App] = ref.get - private[app] def register(app: App): Unit = - ref.getAndSet(Some(app)).foreach { existing => - log.warning( - s"Multiple com.twitter.app.App main methods called. ${existing.name}, then ${app.name}" - ) - } + private[app] def register(app: App): Unit = ref.set(Some(app)) } diff --git a/util-app/src/main/scala/com/twitter/app/ClassPath.scala b/util-app/src/main/scala/com/twitter/app/ClassPath.scala index 308778158c..9dbf20e39c 100644 --- a/util-app/src/main/scala/com/twitter/app/ClassPath.scala +++ b/util-app/src/main/scala/com/twitter/app/ClassPath.scala @@ -1,17 +1,19 @@ package com.twitter.app import com.twitter.finagle.util.loadServiceIgnoredPaths -import java.io.{IOException, File} -import java.net.{URISyntaxException, URLClassLoader, URI} +import java.io.{File, IOException} +import java.net.{URI, URISyntaxException, URL, URLClassLoader} import java.nio.charset.MalformedInputException +import java.nio.file.Paths import java.util.jar.{JarEntry, JarFile} import scala.collection.mutable -import scala.collection.JavaConverters._ +import scala.collection.mutable.Builder import scala.io.Source +import scala.jdk.CollectionConverters._ private[app] object ClassPath { - val IgnoredPackages = Set( + val IgnoredPackages: Set[String] = Set( "apple/", "ch/epfl/", "com/apple/", @@ -51,105 +53,111 @@ private[app] sealed abstract class ClassPath[CpInfo <: ClassPath.Info] { protected def ignoredPackages: Set[String] def browse(loader: ClassLoader): Seq[CpInfo] = { - val buf = mutable.Buffer[CpInfo]() + val buf = Vector.newBuilder[CpInfo] val seenUris = mutable.HashSet[URI]() for ((uri, loader) <- getEntries(loader)) { browseUri0(uri, loader, buf, seenUris) seenUris += uri } - buf + buf.result() + } + + // Note: In JDK 9+ URLClassLoader is no longer the default ClassLoader. + // This method allows us to scan URL's on the class path that can be used + // but does NOT include the ModulePath introduced in JDK9. This method is + // used as a bridge between JDK 8 and JDK 9+. + // The method used here is attributed to https://stackoverflow.com/a/49557901. + // TODO - add suppport for the ModulePath after dropping JDK 8 support. + private[this] def urlsFromClasspath(): Array[URL] = { + val classpath: String = System.getProperty("java.class.path") + classpath.split(File.pathSeparator).map { (pathEntry: String) => + Paths.get(pathEntry).toAbsolutePath().toUri().toURL + } } // package protected for testing private[app] def getEntries(loader: ClassLoader): Seq[(URI, ClassLoader)] = { - val ents = mutable.Buffer[(URI, ClassLoader)]() - val parent = loader.getParent - if (parent != null) - ents ++= getEntries(parent) - - loader match { - case urlLoader: URLClassLoader => - Option(urlLoader.getURLs) match { - case Some(urls) => - urls.foreach { url => - if (url != null) - ents += (url.toURI -> loader) - } - case _ => - } - case _ => - } + val parent = Option(loader.getParent) + + val ownURIs: Vector[(URI, ClassLoader)] = for { + urlLoader <- Vector(loader).map { + case urlClassLoader: URLClassLoader => urlClassLoader + case cl => new URLClassLoader(urlsFromClasspath(), cl) + } + urls <- Option(urlLoader.getURLs()).toVector + url <- urls if url != null + } yield (url.toURI -> loader) + + val p = parent.toSeq.flatMap(getEntries) + + ownURIs ++ p - ents } // package protected for testing private[app] def browseUri( uri: URI, loader: ClassLoader, - buf: mutable.Buffer[CpInfo] + buf: Builder[CpInfo, Seq[CpInfo]] ): Unit = browseUri0(uri, loader, buf, mutable.Set[URI]()) private[this] def browseUri0( uri: URI, loader: ClassLoader, - buf: mutable.Buffer[CpInfo], + buf: Builder[CpInfo, Seq[CpInfo]], history: mutable.Set[URI] ): Unit = { - if (uri.getScheme != "file") - return + if (!history.contains(uri)) { + history.add(uri) + if (uri.getScheme != "file") + return - val f = new File(uri) - if (!(f.exists() && f.canRead)) - return + val f = new File(uri) + if (!(f.exists() && f.canRead)) + return - if (f.isDirectory) - browseDir(f, loader, "", buf) - else - browseJar(f, loader, buf, history) + if (f.isDirectory) + browseDir(f, loader, "", buf) + else + browseJar(f, loader, buf, history) + } } private[this] def browseDir( dir: File, loader: ClassLoader, prefix: String, - buf: mutable.Buffer[CpInfo] + buf: Builder[CpInfo, Seq[CpInfo]] ): Unit = { if (ignoredPackages.contains(prefix)) return for (f <- dir.listFiles) - if (f.isDirectory && f.canRead) + if (f.isDirectory && f.canRead) { browseDir(f, loader, prefix + f.getName + "/", buf) - else + } else processFile(prefix, f, buf) } - protected def processFile( - prefix: String, - file: File, - buf: mutable.Buffer[CpInfo] - ): Unit + protected def processFile(prefix: String, file: File, buf: Builder[CpInfo, Seq[CpInfo]]): Unit private def browseJar( file: File, loader: ClassLoader, - buf: mutable.Buffer[CpInfo], + buf: Builder[CpInfo, Seq[CpInfo]], seenUris: mutable.Set[URI] ): Unit = { - val jarFile = try new JarFile(file) - catch { - case _: IOException => return // not a Jar file - } + val jarFile = + try new JarFile(file) + catch { + case _: IOException => return // not a Jar file + } try { for (uri <- jarClasspath(file, jarFile.getManifest)) { - if (!seenUris.contains(uri)) { - seenUris += uri - browseUri0(uri, loader, buf, seenUris) - } + browseUri0(uri, loader, buf, seenUris) } for { @@ -169,14 +177,14 @@ private[app] sealed abstract class ClassPath[CpInfo <: ClassPath.Info] { protected def processJarEntry( jarFile: JarFile, entry: JarEntry, - buf: mutable.Buffer[CpInfo] + buf: Builder[CpInfo, Seq[CpInfo]] ): Unit private def jarClasspath(jarFile: File, manifest: java.util.jar.Manifest): Seq[URI] = for { m <- Option(manifest).toSeq attr <- Option(m.getMainAttributes.getValue("Class-Path")).toSeq - el <- attr.split(" ") + el <- attr.split(" ").toSeq uri <- uriFromJarClasspath(jarFile, el) } yield uri @@ -204,7 +212,7 @@ private[app] class FlagClassPath extends ClassPath[ClassPath.FlagInfo] { protected def processFile( prefix: String, file: File, - buf: mutable.Buffer[ClassPath.FlagInfo] + buf: Builder[ClassPath.FlagInfo, Seq[ClassPath.FlagInfo]] ): Unit = { val name = file.getName if (isClass(name)) { @@ -215,7 +223,7 @@ private[app] class FlagClassPath extends ClassPath[ClassPath.FlagInfo] { protected def processJarEntry( jarFile: JarFile, entry: JarEntry, - buf: mutable.Buffer[ClassPath.FlagInfo] + buf: Builder[ClassPath.FlagInfo, Seq[ClassPath.FlagInfo]] ): Unit = { val name = entry.getName if (isClass(name)) { @@ -240,7 +248,7 @@ private[app] class LoadServiceClassPath extends ClassPath[ClassPath.LoadServiceI private[app] def readLines(source: Source): Seq[String] = { try { - source.getLines().toArray.flatMap { line => + source.getLines().toVector.flatMap { line => val commentIdx = line.indexOf('#') val end = if (commentIdx != -1) commentIdx else line.length val str = line.substring(0, end).trim @@ -256,7 +264,7 @@ private[app] class LoadServiceClassPath extends ClassPath[ClassPath.LoadServiceI protected def processFile( prefix: String, file: File, - buf: mutable.Buffer[ClassPath.LoadServiceInfo] + buf: Builder[ClassPath.LoadServiceInfo, Seq[ClassPath.LoadServiceInfo]] ): Unit = { for (iface <- ifaceOfName(prefix + file.getName)) { val source = Source.fromFile(file, "UTF-8") @@ -268,7 +276,7 @@ private[app] class LoadServiceClassPath extends ClassPath[ClassPath.LoadServiceI protected def processJarEntry( jarFile: JarFile, entry: JarEntry, - buf: mutable.Buffer[ClassPath.LoadServiceInfo] + buf: Builder[ClassPath.LoadServiceInfo, Seq[ClassPath.LoadServiceInfo]] ): Unit = { for (iface <- ifaceOfName(entry.getName)) { val source = Source.fromInputStream(jarFile.getInputStream(entry), "UTF-8") diff --git a/util-app/src/main/scala/com/twitter/app/Flag.scala b/util-app/src/main/scala/com/twitter/app/Flag.scala index cde12f016d..93f262cb4d 100644 --- a/util-app/src/main/scala/com/twitter/app/Flag.scala +++ b/util-app/src/main/scala/com/twitter/app/Flag.scala @@ -19,7 +19,7 @@ object Flag { * A single command-line flag, instantiated by a [[com.twitter.app.Flags]] * instance. * - * The current value can be extracted via [[apply()]], [[get]], and + * The current value can be extracted via [[apply]], [[get]], and * [[getWithDefault]]. [[com.twitter.util.Local Local]]-scoped modifications * of their values, which can be useful for tests, can be done by calls to * [[let]] and [[letClear]]. @@ -56,15 +56,24 @@ class Flag[T: Flaggable] private[app] ( val name: String, val help: String, defaultOrUsage: Either[() => T, String], - failFastUntilParsed: Boolean -) { + failFastUntilParsed: Boolean) { import com.twitter.app.Flag._ import java.util.logging._ - private[app] def this(name: String, help: String, default: => T, failFastUntilParsed: Boolean) = + private[twitter] def this( + name: String, + help: String, + default: => T, + failFastUntilParsed: Boolean + ) = this(name, help, Left(() => default), failFastUntilParsed) - private[app] def this(name: String, help: String, usage: String, failFastUntilParsed: Boolean) = + private[twitter] def this( + name: String, + help: String, + usage: String, + failFastUntilParsed: Boolean + ) = this(name, help, Right(usage), failFastUntilParsed) private[twitter] def this(name: String, help: String, default: => T) = @@ -73,7 +82,7 @@ class Flag[T: Flaggable] private[app] ( private[twitter] def this(name: String, help: String, usage: String) = this(name, help, usage, false) - protected val flaggable: Flaggable[T] = implicitly[Flaggable[T]] + private[twitter] val flaggable: Flaggable[T] = implicitly[Flaggable[T]] private[this] def localValue: Option[T] = localFlagValues() match { case None => None @@ -147,17 +156,29 @@ class Flag[T: Flaggable] private[app] ( * Override the value of this flag with `t`, only for the scope of the current * [[com.twitter.util.Local]] for the given function `f`. * + * @see [[letParse]] * @see [[letClear]] */ def let[R](t: T)(f: => R): R = let(Some(t), f) + /** + * Override the value of this flag with `arg` `String` parsed to this [[Flag]]'s `T` type, + * only for the scope of the current [[com.twitter.util.Local]] for the given function `f`. + * + * @see [[let]] + * @see [[letClear]] + */ + def letParse[R](arg: String)(f: => R): R = + let(flaggable.parse(arg))(f) + /** * Unset the value of this flag, such that [[isDefined]] will return `false`, * only for the scope of the current [[com.twitter.util.Local]] for the * given function `f`. * * @see [[let]] + * @see [[letParse]] */ def letClear[R](f: => R): R = let(None, f) @@ -211,6 +232,16 @@ class Flag[T: Flaggable] private[app] ( */ def get: Option[T] = getValue + /** + * Get the *unparsed* (as string) value if it has been set. + * + * @note if no user-defined value has been set, `None` will be returned even when a default value + * is supplied. + * + * @see [[Flag.getWithDefaultUnparsed]] + */ + def getUnparsed: Option[String] = getValue.map(flaggable.show) + /** * Get the value if it has been set or if there is a default value supplied. * @@ -219,6 +250,19 @@ class Flag[T: Flaggable] private[app] ( */ def getWithDefault: Option[T] = valueOrDefault + /** + * Get the default value for this flag. This will return `None` if no default value is specified. + */ + private[twitter] def getDefault: Option[T] = default + + /** + * Get the *unparsed* (as string) value if it has been set or if there is a default value supplied. + * + * @see [[Flag.get]] and [[Flag.isDefined]] if you are interested in + * determining if there is a supplied value. + */ + def getWithDefaultUnparsed: Option[String] = valueOrDefault.map(flaggable.show) + /** String representation of this flag's default value */ def defaultString(): String = { try { diff --git a/util-app/src/main/scala/com/twitter/app/Flaggable.scala b/util-app/src/main/scala/com/twitter/app/Flaggable.scala index 1b8ab25833..624b9dfe8a 100644 --- a/util-app/src/main/scala/com/twitter/app/Flaggable.scala +++ b/util-app/src/main/scala/com/twitter/app/Flaggable.scala @@ -1,17 +1,23 @@ package com.twitter.app -import com.twitter.util.{Duration, StorageUnit, Time, TimeFormat} +import com.twitter.util.Duration +import com.twitter.util.StorageUnit +import com.twitter.util.Time +import com.twitter.util.TimeFormat import java.io.File -import java.lang.{ - Boolean => JBoolean, - Double => JDouble, - Float => JFloat, - Integer => JInteger, - Long => JLong -} +import java.lang.reflect.Type +import java.lang.{Boolean => JBoolean} +import java.lang.{Double => JDouble} +import java.lang.{Float => JFloat} +import java.lang.{Integer => JInteger} +import java.lang.{Long => JLong} import java.net.InetSocketAddress -import java.util.{List => JList, Map => JMap, Set => JSet} -import scala.collection.JavaConverters._ +import java.time.LocalTime +import java.util.{List => JList} +import java.util.{Map => JMap} +import java.util.{Set => JSet} +import scala.jdk.CollectionConverters._ +import scala.reflect.ClassTag /** * A type class providing evidence for parsing type `T` as a flag value. @@ -48,7 +54,7 @@ import scala.collection.JavaConverters._ * } * }}} * - * [1] http://en.wikipedia.org/wiki/Type_class + * [1] https://en.wikipedia.org/wiki/Type_class */ abstract class Flaggable[T] { @@ -75,51 +81,73 @@ abstract class Flaggable[T] { */ object Flaggable { + /** + * Just like [[Flaggable]] but provides the runtime with the information about the type `T`. This + * can be useful for frameworks leveraging DI for Twitter Flags. + * + * @note You should prefer this interface for new Flaggables to allow them participate in DI (if + * needed). Eventually this interface may or may not submerge into the parent [[Flaggable]]. + */ + private[twitter] abstract class Typed[T] extends Flaggable[T] { + + /** + * A raw type for this Flaggable. + */ + def rawType: Type + } + + /** + * Just like [[Flaggable.Typed]] but signals the runtime that this Flaggable type is generic so it + * takes types as arguments. For example, `List[T]` is such type and it take one type-parameter + * `T` to construct `List[T]`. + */ + private[twitter] abstract class Generic[T] extends Typed[T] { + + /** + * An [[Array]] of Flaggables that map to type-parameters for this generic Flaggable. + */ + def parameters: Array[Flaggable[_]] + } + /** * Create a `Flaggable[T]` according to a `String => T` conversion function. * * @param f Function that parses a string into an object of type `T` */ - def mandatory[T](f: String => T): Flaggable[T] = new Flaggable[T] { + def mandatory[T](f: String => T)(implicit ct: ClassTag[T]): Flaggable[T] = new Typed[T] { def parse(s: String): T = f(s) + def rawType: Type = ct.runtimeClass } implicit val ofString: Flaggable[String] = mandatory(identity) - // Pairs of Flaggable conversions for types with corresponding Java boxed types. - implicit val ofBoolean: Flaggable[Boolean] = new Flaggable[Boolean] { + implicit val ofBoolean: Flaggable[Boolean] = new Typed[Boolean] { override def default: Option[Boolean] = Some(true) def parse(s: String): Boolean = s.toBoolean + def rawType: Type = classOf[Boolean] } - implicit val ofJavaBoolean: Flaggable[JBoolean] = new Flaggable[JBoolean] { + implicit val ofJavaBoolean: Flaggable[JBoolean] = new Typed[JBoolean] { override def default: Option[JBoolean] = Some(JBoolean.TRUE) def parse(s: String): JBoolean = JBoolean.valueOf(s.toBoolean) + def rawType: Type = classOf[JBoolean] } implicit val ofInt: Flaggable[Int] = mandatory(_.toInt) implicit val ofJavaInteger: Flaggable[JInteger] = - mandatory { s: String => - JInteger.valueOf(s.toInt) - } + mandatory { (s: String) => JInteger.valueOf(s.toInt) } implicit val ofLong: Flaggable[Long] = mandatory(_.toLong) implicit val ofJavaLong: Flaggable[JLong] = - mandatory { s: String => - JLong.valueOf(s.toLong) - } + mandatory { (s: String) => JLong.valueOf(s.toLong) } implicit val ofFloat: Flaggable[Float] = mandatory(_.toFloat) implicit val ofJavaFloat: Flaggable[JFloat] = - mandatory { s: String => - JFloat.valueOf(s.toFloat) - } + mandatory { (s: String) => JFloat.valueOf(s.toFloat) } implicit val ofDouble: Flaggable[Double] = mandatory(_.toDouble) implicit val ofJavaDouble: Flaggable[JDouble] = - mandatory { s: String => - JDouble.valueOf(s.toDouble) - } + mandatory { (s: String) => JDouble.valueOf(s.toDouble) } // Conversions for common non-primitive types and collections. implicit val ofDuration: Flaggable[Duration] = mandatory(Duration.parse(_)) @@ -128,8 +156,10 @@ object Flaggable { private val defaultTimeFormat = new TimeFormat("yyyy-MM-dd HH:mm:ss Z") implicit val ofTime: Flaggable[Time] = mandatory(defaultTimeFormat.parse(_)) + implicit val ofJavaLocalTime: Flaggable[LocalTime] = mandatory(LocalTime.parse) + implicit val ofInetSocketAddress: Flaggable[InetSocketAddress] = - new Flaggable[InetSocketAddress] { + new Typed[InetSocketAddress] { def parse(v: String): InetSocketAddress = v.split(":") match { case Array("", p) => new InetSocketAddress(p.toInt) @@ -140,24 +170,22 @@ object Flaggable { } override def show(addr: InetSocketAddress): String = - "%s:%d".format(Option(addr.getAddress) match { - case Some(a) if a.isAnyLocalAddress => "" - case _ => addr.getHostName - }, addr.getPort) + "%s:%d".format( + Option(addr.getAddress) match { + case Some(a) if a.isAnyLocalAddress => "" + case _ => addr.getHostName + }, + addr.getPort) + + def rawType: Type = classOf[InetSocketAddress] } - implicit val ofFile: Flaggable[File] = new Flaggable[File] { - override def parse(v: String): File = new File(v) - override def show(file: File): String = file.toString - } + implicit val ofFile: Flaggable[File] = mandatory(s => new File(s)) - implicit def ofTuple[T: Flaggable, U: Flaggable]: Flaggable[(T, U)] = new Flaggable[(T, U)] { + implicit def ofTuple[T: Flaggable, U: Flaggable]: Flaggable[(T, U)] = new Generic[(T, U)] { private val tflag = implicitly[Flaggable[T]] private val uflag = implicitly[Flaggable[U]] - assert(tflag.default.isEmpty) - assert(uflag.default.isEmpty) - def parse(v: String): (T, U) = v.split(",") match { case Array(t, u) => (tflag.parse(t), uflag.parse(u)) case _ => throw new IllegalArgumentException("not a 't,u'") @@ -167,29 +195,55 @@ object Flaggable { val (t, u) = tup tflag.show(t) + "," + uflag.show(u) } + + def parameters: Array[Flaggable[_]] = Array(tflag, uflag) + def rawType: Type = classOf[Tuple2[_, _]] + } + + /** + * Create a Flaggable of Java Enum. + * @param clazz The `java.lang.Enum` class. + * @note `java.lang.Enum` enumeration constant value look up is case insensitive. + */ + def ofJavaEnum[T <: Enum[T]](clazz: Class[T]): Flaggable[T] = new Typed[T] { + def parse(s: String): T = + clazz.getEnumConstants.find(_.name.equalsIgnoreCase(s)) match { + case Some(prop) => prop + case _ => + throw new IllegalArgumentException( + s"The property $s does not belong to Java Enum ${clazz.getName}, the constants defined " + + s"in the class are: ${clazz.getEnumConstants.mkString(",")}.") + } + + def rawType: Type = clazz } - private[app] class SetFlaggable[T: Flaggable] extends Flaggable[Set[T]] { + private[app] class SetFlaggable[T: Flaggable] extends Generic[Set[T]] { private val flag = implicitly[Flaggable[T]] - assert(flag.default.isEmpty) - override def parse(v: String): Set[T] = v.split(",").map(flag.parse(_)).toSet + def parse(v: String): Set[T] = + if (v.isEmpty) Set.empty[T] + else v.split(",").iterator.map(flag.parse(_)).toSet override def show(set: Set[T]): String = set.map(flag.show).mkString(",") + + def parameters: Array[Flaggable[_]] = Array(flag) + def rawType: Type = classOf[Set[_]] } - private[app] class SeqFlaggable[T: Flaggable] extends Flaggable[Seq[T]] { + private[app] class SeqFlaggable[T: Flaggable] extends Generic[Seq[T]] { private val flag = implicitly[Flaggable[T]] - assert(flag.default.isEmpty) - def parse(v: String): Seq[T] = v.split(",").map(flag.parse) + def parse(v: String): Seq[T] = + if (v.isEmpty) Seq.empty[T] + else v.split(",").toSeq.map(flag.parse) override def show(seq: Seq[T]): String = seq.map(flag.show).mkString(",") + + override def parameters: Array[Flaggable[_]] = Array(flag) + override def rawType: Type = classOf[Seq[_]] } - private[app] class MapFlaggable[K: Flaggable, V: Flaggable] extends Flaggable[Map[K, V]] { + private[app] class MapFlaggable[K: Flaggable, V: Flaggable] extends Generic[Map[K, V]] { private val kflag = implicitly[Flaggable[K]] private val vflag = implicitly[Flaggable[V]] - assert(kflag.default.isEmpty) - assert(vflag.default.isEmpty) - def parse(in: String): Map[K, V] = { val tuples = in.split(',').foldLeft(Seq.empty[String]) { case (acc, s) if !s.contains('=') => @@ -210,30 +264,40 @@ object Flaggable { override def show(out: Map[K, V]): String = { out.toSeq.map { case (k, v) => k.toString + "=" + v.toString }.mkString(",") } + + def parameters: Array[Flaggable[_]] = Array(kflag, vflag) + def rawType: Type = classOf[Map[_, _]] } implicit def ofSet[T: Flaggable]: Flaggable[Set[T]] = new SetFlaggable[T] implicit def ofSeq[T: Flaggable]: Flaggable[Seq[T]] = new SeqFlaggable[T] implicit def ofMap[K: Flaggable, V: Flaggable]: Flaggable[Map[K, V]] = new MapFlaggable[K, V] - implicit def ofJavaSet[T: Flaggable]: Flaggable[JSet[T]] = new Flaggable[JSet[T]] { + implicit def ofJavaSet[T: Flaggable]: Flaggable[JSet[T]] = new Generic[JSet[T]] { val setFlaggable = new SetFlaggable[T] override def parse(v: String): JSet[T] = setFlaggable.parse(v).asJava override def show(set: JSet[T]): String = setFlaggable.show(set.asScala.toSet) + + def parameters: Array[Flaggable[_]] = setFlaggable.parameters + def rawType: Type = classOf[JSet[_]] } - implicit def ofJavaList[T: Flaggable]: Flaggable[JList[T]] = new Flaggable[JList[T]] { + implicit def ofJavaList[T: Flaggable]: Flaggable[JList[T]] = new Generic[JList[T]] { val seqFlaggable = new SeqFlaggable[T] override def parse(v: String): JList[T] = seqFlaggable.parse(v).asJava - override def show(list: JList[T]): String = seqFlaggable.show(list.asScala) - } + override def show(list: JList[T]): String = seqFlaggable.show(list.asScala.toSeq) - implicit def ofJavaMap[K: Flaggable, V: Flaggable]: Flaggable[JMap[K, V]] = { - val mapFlaggable = new MapFlaggable[K, V] + def parameters: Array[Flaggable[_]] = seqFlaggable.parameters + def rawType: Type = classOf[JList[_]] + } - new Flaggable[JMap[K, V]] { + implicit def ofJavaMap[K: Flaggable, V: Flaggable]: Flaggable[JMap[K, V]] = + new Generic[JMap[K, V]] { + val mapFlaggable = new MapFlaggable[K, V] def parse(in: String): JMap[K, V] = mapFlaggable.parse(in).asJava override def show(out: JMap[K, V]): String = mapFlaggable.show(out.asScala.toMap) + + def parameters: Array[Flaggable[_]] = mapFlaggable.parameters + def rawType: Type = classOf[JMap[_, _]] } - } } diff --git a/util-app/src/main/scala/com/twitter/app/Flags.scala b/util-app/src/main/scala/com/twitter/app/Flags.scala index 87022b8fe3..a2c7c32bc5 100644 --- a/util-app/src/main/scala/com/twitter/app/Flags.scala +++ b/util-app/src/main/scala/com/twitter/app/Flags.scala @@ -3,6 +3,8 @@ package com.twitter.app import scala.collection.immutable.TreeSet import scala.collection.mutable import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ +import scala.reflect.ClassTag import scala.util.control.NonFatal /** @@ -11,7 +13,7 @@ import scala.util.control.NonFatal * account of incorrectly-interpreted or malformed flag values. * * @param message A string name of the flag for which parsing failed. - * @param cause The underlying [[java.lang.Throwable]] that caused this exception. + * @param cause The underlying `java.lang.Throwable` that caused this exception. */ case class FlagParseException(message: String, cause: Throwable = null) extends Exception(message, cause) @@ -24,7 +26,7 @@ class FlagUndefinedException extends Exception(Flags.FlagUndefinedMessage) object Flags { private[app] val FlagValueRequiredMessage = "flag value is required" - private[app] val FlagUndefinedMessage = "flag undefined" + private[app] val FlagUndefinedMessage = "flag instance (c.t.app.Flag) is not defined" sealed trait FlagParseResult @@ -47,7 +49,9 @@ object Flags { * * @param reason A string explaining the error that occurred. */ - case class Error(reason: String) extends FlagParseResult + case class Error(reason: String) extends FlagParseResult { + override def toString: String = reason + } } /** @@ -79,7 +83,7 @@ object Flags { * be included during flag parsing. If false, only flags defined in the * application itself will be consulted. */ -class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) { +final class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) { import com.twitter.app.Flags._ def this(argv0: String, includeGlobal: Boolean) = this(argv0, includeGlobal, false) @@ -87,6 +91,8 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) // thread-safety is provided by synchronization on `this` private[this] val flags = new mutable.HashMap[String, Flag[_]] + // when a flag with the same name as a previously added flag is added, track it here + private[this] val duplicates = new java.util.concurrent.ConcurrentSkipListSet[String]() @volatile private[this] var cmdUsage = "" @@ -120,83 +126,84 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @return A [[com.twitter.app.Flags.FlagParseResult]] representing the * result of parsing `args`. */ - def parseArgs( - args: Array[String], - allowUndefinedFlags: Boolean = false - ): FlagParseResult = synchronized { - reset() - val remaining = new ArrayBuffer[String] - var i = 0 - while (i < args.length) { - val a = args(i) - i += 1 - if (a == "--") { - remaining ++= args.slice(i, args.length) - i = args.length - } else if (a startsWith "-") { - a drop 1 split ("=", 2) match { - // There seems to be a bug Scala's pattern matching - // optimizer that leaves `v' dangling in the last case if - // we make this a wildcard (Array(k, _@_*)) - case Array(k) if !hasFlag(k) => - if (allowUndefinedFlags) - remaining += a - else - return Error( - "Error parsing flag \"%s\": %s".format(k, FlagUndefinedMessage) - ) + def parseArgs(args: Array[String], allowUndefinedFlags: Boolean = false): FlagParseResult = + synchronized { + reset() + val remaining = new ArrayBuffer[String] + val errors = new ArrayBuffer[Error] + var i = 0 + while (i < args.length) { + val a = args(i) + i += 1 + if (a == "--") { + remaining ++= args.slice(i, args.length) + i = args.length + } else if (a startsWith "-") { + a drop 1 split ("=", 2) match { + // There seems to be a bug Scala's pattern matching + // optimizer that leaves `v' dangling in the last case if + // we make this a wildcard (Array(k, _@_*)) + case Array(k) if !hasFlag(k) => + if (allowUndefinedFlags) + remaining += a + else + errors += Error( + "Error parsing flag \"%s\": %s".format(k, FlagUndefinedMessage) + ) + + // Flag isn't defined + case Array(k, _) if !hasFlag(k) => + if (allowUndefinedFlags) + remaining += a + else + errors += Error( + "Error parsing flag \"%s\": %s".format(k, FlagUndefinedMessage) + ) - // Flag isn't defined - case Array(k, _) if !hasFlag(k) => - if (allowUndefinedFlags) - remaining += a - else - return Error( - "Error parsing flag \"%s\": %s".format(k, FlagUndefinedMessage) + // Optional argument without a value + case Array(k) if flag(k).noArgumentOk => + flag(k).parse() + + // Mandatory argument without a value and with no more arguments. + case Array(k) if i == args.length => + errors += Error( + "Error parsing flag \"%s\": %s".format(k, FlagValueRequiredMessage) ) - // Optional argument without a value - case Array(k) if flag(k).noArgumentOk => - flag(k).parse() - - // Mandatory argument without a value and with no more arguments. - case Array(k) if i == args.length => - return Error( - "Error parsing flag \"%s\": %s".format(k, FlagValueRequiredMessage) - ) - - // Mandatory argument with another argument - case Array(k) => - i += 1 - try flag(k).parse(args(i - 1)) - catch { - case NonFatal(e) => - return Error( - "Error parsing flag \"%s\": %s".format(k, e.getMessage) - ) - } - - // Mandatory k=v - case Array(k, v) => - try flag(k).parse(v) - catch { - case e: Throwable => - return Error( - "Error parsing flag \"%s\": %s".format(k, e.getMessage) - ) - } + // Mandatory argument with another argument + case Array(k) => + i += 1 + try flag(k).parse(args(i - 1)) + catch { + case NonFatal(e) => + errors += Error( + "Error parsing flag \"%s\": %s".format(k, e.getMessage) + ) + } + + // Mandatory k=v + case Array(k, v) => + try flag(k).parse(v) + catch { + case NonFatal(e) => + errors += Error( + "Error parsing flag \"%s\": %s".format(k, e.getMessage) + ) + } + } + } else { + remaining += a } - } else { - remaining += a } + finishParsing() + + if (helpFlag()) + Help(usage) + else if (errors.nonEmpty) + Error(s"Error parsing flags: ${errors.mkString("[\n ", ",\n ", "\n]")}\n\n$usage") + else + Ok(remaining.toSeq) } - finishParsing() - - if (helpFlag()) - Help(usage) - else - Ok(remaining) - } /** * Parse an array of flag strings. @@ -212,14 +219,12 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @throws FlagParseException if an error occurs during flag-parsing. */ @deprecated("Prefer result-value based `Flags.parseArgs` method", "6.17.1") - def parse( - args: Array[String], - undefOk: Boolean = false - ): Seq[String] = parseArgs(args, undefOk) match { - case Ok(remainder) => remainder - case Help(usage) => throw FlagUsageError(usage) - case Error(reason) => throw FlagParseException(reason) - } + def parse(args: Array[String], undefOk: Boolean = false): Seq[String] = + parseArgs(args, undefOk) match { + case Ok(remainder) => remainder + case Help(usage) => throw FlagUsageError(usage) + case Error(reason) => throw FlagParseException(reason) + } /** * Parse an array of flag strings or exit the application (with exit code 1) @@ -252,8 +257,10 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @param f A concrete Flag to add */ def add(f: Flag[_]): Unit = synchronized { - if (flags contains f.name) - System.err.printf("Flag %s already defined!\n", f.name) + if (flags.contains(f.name)) { + System.err.printf("Flag: \"%s\" already defined!\n", f.name) + duplicates.add(f.name) + } flags(f.name) = f.withFailFast(failFastUntilParsed) } @@ -276,7 +283,7 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @param name The name of the flag. * @param help The help string of the flag. */ - def apply[T](name: String, help: String)(implicit _f: Flaggable[T], m: Manifest[T]): Flag[T] = { + def apply[T](name: String, help: String)(implicit _f: Flaggable[T], m: ClassTag[T]): Flag[T] = { val f = new Flag[T](name, help, m.toString, failFastUntilParsed) add(f) f @@ -290,7 +297,7 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @param help The help string of the flag. */ def create[T](name: String, default: T, help: String, flaggable: Flaggable[T]): Flag[T] = { - implicit val impl = flaggable + implicit val impl: Flaggable[T] = flaggable apply(name, default, help) } @@ -335,13 +342,9 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) cmd + argv0 + " [...]\n" + "flags:\n" + - (lines mkString "\n") + ( - if (globalLines.isEmpty) "" - else { - "\nglobal flags:\n" + - (globalLines mkString "\n") - } - ) + lines.mkString("\n") + + (if (globalLines.isEmpty) "" + else "\nglobal flags:\n" + globalLines.mkString("\n")) } /** @@ -350,7 +353,7 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) * @param includeGlobal If true, all registered * [[com.twitter.app.GlobalFlag GlobalFlags]] will be included * in output. Defaults to the `includeGlobal` settings of this instance. - * @param classLoader The [[java.lang.ClassLoader]] used to fetch + * @param classLoader The `java.lang.ClassLoader` used to fetch * [[com.twitter.app.GlobalFlag GlobalFlags]], if necessary. Defaults to this * instance's classloader. * @return All of the flags known to this this Flags instance. @@ -382,7 +385,7 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) ): (Iterable[String], Iterable[String]) = { val (set, unset) = getAll(includeGlobal, classLoader).partition { _.get.isDefined } - (set.map { _ + " \\" }, unset.map { _ + " \\" }) + (set.map { _.toString + " \\" }, unset.map { _.toString + " \\" }) } /** @@ -405,4 +408,8 @@ class Flags(argv0: String, includeGlobal: Boolean, failFastUntilParsed: Boolean) lines.mkString("\n") } + + /** Return the set of any registered duplicated flag names. */ + protected[twitter] def registeredDuplicates: Set[String] = + duplicates.asScala.toSet } diff --git a/util-app/src/main/scala/com/twitter/app/GlobalFlag.scala b/util-app/src/main/scala/com/twitter/app/GlobalFlag.scala index 8fb6605d01..bc57d68ab5 100644 --- a/util-app/src/main/scala/com/twitter/app/GlobalFlag.scala +++ b/util-app/src/main/scala/com/twitter/app/GlobalFlag.scala @@ -1,8 +1,8 @@ package com.twitter.app -import java.lang.reflect.{Method, Modifier} -import scala.collection.mutable.ArrayBuffer +import java.lang.reflect.Modifier import scala.util.control.NonFatal +import scala.reflect.ClassTag /** * Subclasses of GlobalFlag (that are defined in libraries) are "global" in the @@ -44,9 +44,17 @@ import scala.util.control.NonFatal * If you'd like to declare a new [[GlobalFlag]] in Java, see [[JavaGlobalFlag]]. */ @GlobalFlagVisible -abstract class GlobalFlag[T] private[app] (defaultOrUsage: Either[() => T, String], help: String)( - implicit _f: Flaggable[T] -) extends Flag[T](null, help, defaultOrUsage, false) { +abstract class GlobalFlag[T] private[app] ( + defaultOrUsage: Either[() => T, String], + help: String +)( + implicit _f: Flaggable[T]) + extends Flag[T](null, help, defaultOrUsage, false) { + + // Flags defined inside package objects have unexpected names. + require( + !GlobalFlag.isEnclosedInPackageObject(getClass), + s"package object encloses flag definition: $getClass") override protected[this] def parsingDone: Boolean = true @@ -78,7 +86,7 @@ abstract class GlobalFlag[T] private[app] (defaultOrUsage: Either[() => T, Strin * * @param help documentation regarding usage of this [[Flag]]. */ - def this(help: String)(implicit _f: Flaggable[T], m: Manifest[T]) = + def this(help: String)(implicit _f: Flaggable[T], m: ClassTag[T]) = this(Right(m.toString), help) /** @@ -113,29 +121,44 @@ abstract class GlobalFlag[T] private[app] (defaultOrUsage: Either[() => T, Strin object GlobalFlag { - private[app] def get(f: String): Option[Flag[_]] = { - def validMethod(m: Method): Boolean = - m != null && - Modifier.isStatic(m.getModifiers) && - m.getReturnType == classOf[Flag[_]] && - m.getParameterCount == 0 + private def isEnclosedInPackageObject(klass: Class[_]): Boolean = + klass.getEnclosingClass != null && + klass.getEnclosingClass.getSimpleName == "package" - def tryMethod(clsName: String, methodName: String): Option[Flag[_]] = + private[app] def get(flagName: String): Option[Flag[_]] = { + val className = if (!flagName.endsWith("$")) flagName + "$" else flagName + try { + get(Class.forName(className)) + } catch { case _: ClassNotFoundException => None } + } + + private[app] def get(cls: Class[_]): Option[Flag[_]] = { + def tryMethod(cls: Class[_], methodName: String): Option[Flag[_]] = try { - val cls = Class.forName(clsName) val m = cls.getMethod(methodName) - if (validMethod(m)) - Some(m.invoke(null).asInstanceOf[Flag[_]]) - else + val isValid = Modifier.isStatic(m.getModifiers) && + m.getReturnType == classOf[Flag[_]] && + m.getParameterCount == 0 + if (isValid) Some(m.invoke(null).asInstanceOf[Flag[_]]) else None + } catch { + case _: NoSuchMethodException | _: IllegalArgumentException => None + } + + def tryModuleField(cls: Class[_]): Option[Flag[_]] = + try { + val f = cls.getField("MODULE$") + val isValid = Modifier.isStatic(f.getModifiers) && classOf[Flag[_]] + .isAssignableFrom(f.getType) + if (isValid) Some(f.get(null).asInstanceOf[Flag[_]]) else None } catch { - case _: ClassNotFoundException | _: NoSuchMethodException | _: IllegalArgumentException => + case _: NoSuchFieldException | _: IllegalArgumentException => None } - tryMethod(f, "getGlobalFlag").orElse { + tryModuleField(cls).orElse { // fallback for GlobalFlags declared in Java - tryMethod(f + "$", "globalFlagInstance") + tryMethod(cls, "globalFlagInstance") } } @@ -149,7 +172,7 @@ object GlobalFlag { //we don't hide the real issue that a developer just added an unparseable arg. case e: Throwable => log.log(java.util.logging.Level.SEVERE, "failure reading in flags", e) - new ArrayBuffer[Flag[_]] + Nil } } @@ -162,7 +185,7 @@ object GlobalFlag { className.endsWith("$") && !className.endsWith("package$") val markerClass = classOf[GlobalFlagVisible] - val flags = new ArrayBuffer[Flag[_]] + val flagBuilder = new scala.collection.immutable.VectorBuilder[Flag[_]] // Search for Scala objects annotated with GlobalFlagVisible: val cp = new FlagClassPath() @@ -170,15 +193,15 @@ object GlobalFlag { try { val cls: Class[_] = Class.forName(info.className, false, loader) if (cls.isAnnotationPresent(markerClass)) { - get(info.className.dropRight(1)) match { - case Some(f) => flags += f - case None => println("failed for " + info.className) + get(cls) match { + case Some(f) => flagBuilder += f + case None => log.log(java.util.logging.Level.INFO, "failed for " + info.className) } } } catch { case _: IllegalStateException | _: NoClassDefFoundError | _: ClassNotFoundException => } } - flags + flagBuilder.result() } } diff --git a/util-app/src/main/scala/com/twitter/app/JavaGlobalFlag.scala b/util-app/src/main/scala/com/twitter/app/JavaGlobalFlag.scala index b4df2c38fb..28b89102d8 100644 --- a/util-app/src/main/scala/com/twitter/app/JavaGlobalFlag.scala +++ b/util-app/src/main/scala/com/twitter/app/JavaGlobalFlag.scala @@ -1,21 +1,19 @@ package com.twitter.app /** - * @define java_declaration - * * In order to declare a new [[GlobalFlag]] in Java, it can be done with a bit * of ceremony by following a few steps: * - * 1. the flag's class name must end with "$", e.g. `javaGlobalWithDefault$`. - * This is done to mimic the behavior of Scala's singleton `object`. - * 1. the class must have a method, `public static Flag globalFlagInstance()`, - * that returns the singleton instance of the [[Flag]]. - * 1. the class should have a `static final` field holding the singleton - * instance of the flag. - * 1. the class should extend `JavaGlobalFlag`. - * 1. to avoid extraneous creation of instances, thus forcing users to only - * access the singleton, the constructor should be `private` and the class - * should be `public final`. + * 1. the flag's class name must end with "$", e.g. `javaGlobalWithDefault$`. + * This is done to mimic the behavior of Scala's singleton `object`. + * 1. the class must have a method, `public static Flag globalFlagInstance()`, + * that returns the singleton instance of the [[Flag]]. + * 1. the class should have a `static final` field holding the singleton + * instance of the flag. + * 1. the class should extend `JavaGlobalFlag`. + * 1. to avoid extraneous creation of instances, thus forcing users to only + * access the singleton, the constructor should be `private` and the class + * should be `public final`. * * An example of a [[GlobalFlag]] declared in Java: * {{{ @@ -37,8 +35,8 @@ package com.twitter.app abstract class JavaGlobalFlag[T] private ( defaultOrUsage: Either[() => T, String], help: String, - flaggable: Flaggable[T] -) extends GlobalFlag[T](defaultOrUsage, help)(flaggable) { + flaggable: Flaggable[T]) + extends GlobalFlag[T](defaultOrUsage, help)(flaggable) { /** * A flag with a default value. diff --git a/util-app/src/main/scala/com/twitter/app/Lifecycle.scala b/util-app/src/main/scala/com/twitter/app/Lifecycle.scala new file mode 100644 index 0000000000..84be72f680 --- /dev/null +++ b/util-app/src/main/scala/com/twitter/app/Lifecycle.scala @@ -0,0 +1,32 @@ +package com.twitter.app + +import com.twitter.app.lifecycle.{Event, Notifier, Observer} +import com.twitter.util.Future +import java.util.concurrent.ConcurrentLinkedQueue +import scala.jdk.CollectionConverters._ + +/** + * Mixin to an App that allows for observing lifecycle events. + * To emit an event + * {{{ observe(PreMain) { // logic for premain } }}} + */ +private[app] trait Lifecycle { self: App => + private[this] val observers = new ConcurrentLinkedQueue[Observer]() + private[this] val notifier = new Notifier(observers.asScala) + + /** + * Notifies `Observer`s of `Event`s. + */ + protected final def observe(event: Event)(f: => Unit): Unit = + notifier(event)(f) + + /** + * Notifies `Observer`s of [[com.twitter.util.Future]] `Event`s. + */ + protected final def observeFuture(event: Event)(f: Future[Unit]): Future[Unit] = + notifier.future(event)(f) + + /** Add a lifecycle [[com.twitter.app.lifecycle.Event]] [[com.twitter.app.lifecycle.Observer]] */ + private[twitter] final def withObserver(observer: Observer): Unit = + observers.add(observer) +} diff --git a/util-app/src/main/scala/com/twitter/app/LoadService.scala b/util-app/src/main/scala/com/twitter/app/LoadService.scala index 6ee368356b..e80a0fe66c 100644 --- a/util-app/src/main/scala/com/twitter/app/LoadService.scala +++ b/util-app/src/main/scala/com/twitter/app/LoadService.scala @@ -6,18 +6,18 @@ import java.util.ServiceConfigurationError import java.util.concurrent.ConcurrentHashMap import java.util.function.{Function => JFunction} import java.util.logging.{Level, Logger} -import scala.collection.JavaConverters._ import scala.collection.mutable import scala.io.Source +import scala.jdk.CollectionConverters._ import scala.reflect.ClassTag import scala.util.control.NonFatal /** - * Load classes in the manner of [[java.util.ServiceLoader]]. It is + * Load classes in the manner of `java.util.ServiceLoader`. It is * more resilient to varying Java packaging configurations than ServiceLoader. * * For further control, see [[App.loadServiceBindings]] and the - * [[loadServiceDenied]] flag. + * `loadServiceDenied` flag. */ object LoadService { @@ -28,10 +28,8 @@ object LoadService { * @param implementations cannot be empty. * @see [[App.loadServiceBindings]] */ - class Binding[T]( - val iface: Class[T], - val implementations: Seq[T] - ) { + class Binding[T](val iface: Class[T], val implementations: Seq[T]) { + /** * Convenience constructor for binding a single implementation. */ @@ -43,7 +41,7 @@ object LoadService { * @param implementations cannot be empty. */ def this(iface: Class[T], implementations: java.util.List[T]) = - this(iface, implementations.asScala) + this(iface, implementations.asScala.toSeq) if (implementations.isEmpty) { throw new IllegalArgumentException("implementations cannot be empty") @@ -82,11 +80,11 @@ object LoadService { * * The public API entry point for this is [[App.loadServiceBindings]]. * - * If [[bind]] is called for a `Class[T]` after [[LoadService.apply]] + * If `bind` is called for a `Class[T]` after [[LoadService.apply]] * has been called for that same interface, an `IllegalStateException` * will be thrown. * - * If [[bind]] is called multiple times for the same interface, the + * If `bind` is called multiple times for the same interface, the * last one registered is used, but warnings and lint violations will * be emitted. * @@ -100,19 +98,22 @@ object LoadService { val prevs = binds.synchronized { if (loaded.contains(binding.iface)) { throw new IllegalStateException( - s"LoadService.apply has already been called for ${binding.iface.getName}.") + s"LoadService.apply has already been called for ${binding.iface.getName}." + ) } binds.put(binding.iface, binding.implementations) } prevs.foreach { ps => bindDupes.put(binding.iface, ()) - log.warning(s"Replaced existing `LoadService.bind` for `${binding.iface.getName}`, " + - s"old='$ps', new='${binding.implementations}'") + log.warning( + s"Replaced existing `LoadService.bind` for `${binding.iface.getName}`, " + + s"old='$ps', new='${binding.implementations}'" + ) } } /** - * Return any interfaces that have been registered multiple times via [[bind]]. + * Return any interfaces that have been registered multiple times via `bind`. */ def duplicateBindings: Set[Class[_]] = bindDupes.keySet.asScala.toSet @@ -140,8 +141,6 @@ object LoadService { /** * Returns classes for the given `ClassTag` as specified by * resource files in `META-INF/services/FullyQualifiedClassTagsClassName`. - * - * @see [[apply(Class)]] for a Java friendly API. */ def apply[T: ClassTag](): Seq[T] = { val iface = implicitly[ClassTag[T]].runtimeClass.asInstanceOf[Class[T]] @@ -156,7 +155,7 @@ object LoadService { val cp = new LoadServiceClassPath() val whenAbsent = new JFunction[ClassLoader, Seq[ClassPath.LoadServiceInfo]] { def apply(loader: ClassLoader): Seq[ClassPath.LoadServiceInfo] = { - cp.browse(loader) + cp.browse(loader).toSeq } } diff --git a/util-app/src/main/scala/com/twitter/app/command/Command.scala b/util-app/src/main/scala/com/twitter/app/command/Command.scala new file mode 100644 index 0000000000..1ed0eae60a --- /dev/null +++ b/util-app/src/main/scala/com/twitter/app/command/Command.scala @@ -0,0 +1,51 @@ +package com.twitter.app.command + +import java.io.File + +/** + * A utility to allow for processing the outout from an external process in a streaming manner. + * When [[Command.run]] is called, a [[CommandOutput]] is returned, which provides a + * [[com.twitter.io.Reader]] interface to both the stdout and stderr of the external process. + * + * Example usage: + * {{{ + * // Run a shell command and log the output as it shows up, and then discards the output + * + * val commandOutput: CommandOutput = Command.run(Seq("./slow-shell-command.sh") + * val logged: Reader[Unit] = commandOutput.stdout.map { + * case Buf.Utf8(str) => info(s"stdout: $str") + * } + * Await.result(Reader.readAllLines(logged).onFailure { + * case NonZeroReturnCode(_, code, stderr) => + * val errors = stderr.map { + * case Buf.Utf8(str) => str + * } + * error(s"failure: $code") + * errors.foreach(errLine => error(s"stderr: $errLine") + * }) + * }}} + */ +object Command extends Command(new CommandExecutorImpl) + +class Command private[command] ( + commandExecutor: CommandExecutor) { + + /** + * Runs a shell command and returns a [[CommandOutput]]. See the CommandOutput scaladoc + * for more details on the behavior of fields. + * + * @param cmd The shell command to run + * @param workingDirectory An optional directory from which the command should be run + * @param extraEnv A map of extra environment variables to set for the command to be run + * @return A [[CommandOutput]] with stdout & stderr from the command + */ + def run( + cmd: Seq[String], + workingDirectory: Option[File] = None, + extraEnv: Map[String, String] = Map.empty + ): CommandOutput = { + val process = commandExecutor(cmd, workingDirectory, extraEnv) + CommandOutput.fromProcess(cmd, process) + } + +} diff --git a/util-app/src/main/scala/com/twitter/app/command/CommandExecutor.scala b/util-app/src/main/scala/com/twitter/app/command/CommandExecutor.scala new file mode 100644 index 0000000000..ee6c91cab7 --- /dev/null +++ b/util-app/src/main/scala/com/twitter/app/command/CommandExecutor.scala @@ -0,0 +1,37 @@ +package com.twitter.app.command + +import java.io.File +import scala.jdk.CollectionConverters._ + +/** + * CommandExecutor is a private trait used for testing so that the actual forking of the + * command can be mocked out. + */ +private[command] trait CommandExecutor { + def apply( + cmd: Seq[String], + workingDirectory: Option[File], + extraEnv: Map[String, String] + ): Process +} + +private[command] class CommandExecutorImpl extends CommandExecutor { + + /** + * + * @param cmd A [[Seq]] of [[String]] for the actual command + * @param workingDirectory The directory in which the command will be run + * @param extraEnv A map of extra environment variables to set for the command to be run + * @return a [[java.lang.Process]] representing the subprocess kicked off + */ + def apply( + cmd: Seq[String], + workingDirectory: Option[File], + extraEnv: Map[String, String] + ): Process = { + val processBuilder = new java.lang.ProcessBuilder(cmd: _*) + workingDirectory.foreach(dir => processBuilder.directory(dir)) + processBuilder.environment().putAll(extraEnv.asJava) + processBuilder.start() + } +} diff --git a/util-app/src/main/scala/com/twitter/app/command/CommandOutput.scala b/util-app/src/main/scala/com/twitter/app/command/CommandOutput.scala new file mode 100644 index 0000000000..f9a3a74635 --- /dev/null +++ b/util-app/src/main/scala/com/twitter/app/command/CommandOutput.scala @@ -0,0 +1,95 @@ +package com.twitter.app.command + +import com.twitter.io.{Buf, Pipe, Reader} +import com.twitter.util.{Future, FuturePool} +import scala.util.control.NonFatal + +object CommandOutput { + private val Delim = '\n'.toByte + private val DelimLength = 1 + private val futurePool = FuturePool.unboundedPool + + /** + * Turns a [[java.lang.Process]] into a [[CommandOutput]]. The stdout and stderr inputstreams are + * converted into a newline-delimited [[Reader]]. If the [[Process]] exits with a non-zero + * return code, the final [[Reader#read]] on stdout will return an exceptional [[Future]]. + * + * @param cmd The command used to create the process (used for the exception message) + * @param process The [[java.lang.Process]] that we are turning into a CommandOutput + * @return + */ + def fromProcess(cmd: Seq[String], process: Process): CommandOutput = { + val outputStream = splitOnNewLines(Reader.fromStream(process.getInputStream)) + val errorStream = splitOnNewLines(Reader.fromStream(process.getErrorStream)) + + val returnCode = futurePool { + process.waitFor() + } + + val endReader = Reader + .fromFuture(returnCode.map { + case 0 => + Reader.empty + case failure => + Reader.exception(NonZeroReturnCode(cmd, failure, errorStream)) + }).flatten + + CommandOutput(Reader.concat(Seq(outputStream, endReader)), errorStream) + } + + /** + * Takes a [[com.twitter.io.Reader]] of [[com.twitter.io.Buf]] and transforms to a [[Reader]] of + * [[Buf]] where each element in the resulting reader is the original data delimited by + * a newline. Delimiters are not included in the output. + * + * @param in The incoming [[Reader]] + * @return A [[Reader]] of the delimited records + * + */ + private[command] def splitOnNewLines( + in: Reader[Buf] + ): Reader[Buf] = { + val pipe = new Pipe[Buf]() + + def copyLoop(accum: Buf): Future[Unit] = { + val delimIndex: Int = accum.process((byte: Byte) => byte != Delim) + if (delimIndex != -1) { + pipe.write(accum.slice(0, delimIndex)).before { + copyLoop(accum.slice(delimIndex + DelimLength, Int.MaxValue)) + } + } else { + in.read().flatMap { + case Some(buf) => copyLoop(accum.concat(buf)) + case None => + (if (!accum.isEmpty) pipe.write(accum) else Future.Unit).before { + pipe.close() + } + }.rescue { + case NonFatal(e) => + pipe.fail(e) + Future.exception(e) + } + } + } + copyLoop(Buf.Empty).onFailure { + case NonFatal(_) => + in.discard() + } + pipe + } +} + +/** + * The output of a shell command. + * + * @param stdout A [[com.twitter.io.Reader]] representing what is written to stdout + * by the shell command. Each [[Buf]] is a line of output from + * the command. If the command exits with a non zero return code, + * the final [[Reader#read()]] will return a failed [[com.twitter.util.Future]] + * with a [[NonZeroReturnCode]] exception. If the shell command + * completes successfully the final read will return a [[None]] + * @param stderr A [[Reader]] with the stderr from the shell command. The final + * read on stderr will always return a [[None]], even if the command exited with a + * nonzero return code. + */ +case class CommandOutput(stdout: Reader[Buf], stderr: Reader[Buf]) diff --git a/util-app/src/main/scala/com/twitter/app/command/NonZeroReturnCode.scala b/util-app/src/main/scala/com/twitter/app/command/NonZeroReturnCode.scala new file mode 100644 index 0000000000..4293b17d4a --- /dev/null +++ b/util-app/src/main/scala/com/twitter/app/command/NonZeroReturnCode.scala @@ -0,0 +1,14 @@ +package com.twitter.app.command + +import com.twitter.io.{Buf, Reader} + +/** + * The exception thrown when the command returns with a nonzero return code. + * + * @param code The actual return code from the command + * @param stdErr A [[Reader]] with the stderr from the command. Note that if the stderr has already + * been consumed, then this [[Reader]] will only contain the portion of the stderr, + * that has not yet been consumed. + */ +case class NonZeroReturnCode(command: Seq[String], code: Int, stdErr: Reader[Buf]) + extends Exception(s"Command $command exited with return code: $code") diff --git a/util-app/src/main/scala/com/twitter/finagle/util/flags.scala b/util-app/src/main/scala/com/twitter/finagle/util/flags.scala index 62a4d60387..8bab4a7f38 100644 --- a/util-app/src/main/scala/com/twitter/finagle/util/flags.scala +++ b/util-app/src/main/scala/com/twitter/finagle/util/flags.scala @@ -1,6 +1,7 @@ package com.twitter.finagle.util import com.twitter.app +import com.twitter.app.GlobalFlag /** * A deny list of implementations to ignore. Keys are the fully qualified class names. @@ -44,3 +45,12 @@ object loadServiceIgnoredPaths Seq.empty[String], "Additional packages to be excluded from recursive directory scan" ) + +/** + * This flag enables another way for graceful shutdown. + */ +object enableJvmShutdownHook + extends GlobalFlag[Boolean]( + true, + "Registers a JVM shutdown hook which invokes close()" + ) diff --git a/util-app/src/test/java/BUILD b/util-app/src/test/java/BUILD index a1f3950237..54e9a51b22 100644 --- a/util-app/src/test/java/BUILD +++ b/util-app/src/test/java/BUILD @@ -1,12 +1,19 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-app/src/main/java", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-app/src/main/java/com/twitter/app", "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-function/src/main/java", + "util/util-core:util-core-util", + "util/util-function/src/main/java/com/twitter/function", ], ) diff --git a/util-app/src/test/java/com/twitter/app/JavaFlagTest.java b/util-app/src/test/java/com/twitter/app/JavaFlagTest.java index 8fab172321..38d8fe5ac4 100644 --- a/util-app/src/test/java/com/twitter/app/JavaFlagTest.java +++ b/util-app/src/test/java/com/twitter/app/JavaFlagTest.java @@ -6,6 +6,8 @@ import java.util.List; import java.util.Map; +import scala.reflect.ClassTag$; + import org.junit.Test; import com.twitter.util.Duration; @@ -32,13 +34,18 @@ static class Named { @Override public Named apply(String n) { return new Named(n); } - }); + }, ClassTag$.MODULE$.apply(Named.class)); Named(String name) { this.name = name; } } + enum Disney { + MICKEY, + MINNIE + } + @Test public void testJavaFlags() { @@ -98,6 +105,33 @@ public void testJavaLongFlags() { } } + @Test + public void testJavaEnumFlags() { + Flags flags = new Flags("ApplicationName"); + Flag enumFlag = flags.create( + "disney", + Disney.MICKEY, + "", + Flaggable.ofJavaEnum(Disney.class) + ); + enumFlag.parse(); + assertEquals(Disney.MICKEY, enumFlag.apply()); + // should parse Strings with matching cases without errors + enumFlag.parse("MICKEY"); + // should parse Strings without matching cases without errors + enumFlag.parse("minnie"); + // should thrown IllegalArgumentException when property doesn't exist in the Enum + try{ + enumFlag.parse("MICKY"); + } catch (IllegalArgumentException e) { + assertEquals( + "The property MICKY does not belong to Java Enum com.twitter.app.JavaFlagTest$Disney, " + + "the constants defined in the class are: MICKEY,MINNIE.", + e.getMessage() + ); + } + } + @Test public void testLet() { Flags flags = new Flags("ApplicationName"); diff --git a/util-app/src/test/resources/BUILD b/util-app/src/test/resources/BUILD index e0afa86adb..07015f79b6 100644 --- a/util-app/src/test/resources/BUILD +++ b/util-app/src/test/resources/BUILD @@ -1,6 +1,8 @@ resources( - sources = rglobs( - "*", - exclude = [rglobs("BUILD*", "*.pyc")], - ), + sources = [ + "!**/*.pyc", + "!BUILD*", + "**/*", + ], + tags = ["bazel-compatible"], ) diff --git a/util-app/src/test/resources/command-test.sh b/util-app/src/test/resources/command-test.sh new file mode 100755 index 0000000000..525fddc53f --- /dev/null +++ b/util-app/src/test/resources/command-test.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash +# +# Simple script used for testing Command +# +REPS=$1 +TIME=$2 +EXITCODE=$3 + +for rep in $(seq 1 "$REPS"); do + echo "Stdout # $rep $EXTRA_ENV" + >&2 echo "Stderr # $rep $EXTRA_ENV" + sleep "$TIME" +done + +exit "$EXITCODE" diff --git a/util-app/src/test/scala/BUILD b/util-app/src/test/scala/BUILD index 67a9019c1c..e26138ec02 100644 --- a/util-app/src/test/scala/BUILD +++ b/util-app/src/test/scala/BUILD @@ -1,13 +1,22 @@ junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/mockito:mockito-core", "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", "util/util-app/src/main/scala", "util/util-app/src/test/resources", - "util/util-core/src/main/scala", - "util/util-registry/src/main/scala", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-registry/src/main/scala/com/twitter/util/registry", ], ) diff --git a/util-app/src/test/scala/com/twitter/app/AppTest.scala b/util-app/src/test/scala/com/twitter/app/AppTest.scala index 0e38ecfa58..f86edb7770 100644 --- a/util-app/src/test/scala/com/twitter/app/AppTest.scala +++ b/util-app/src/test/scala/com/twitter/app/AppTest.scala @@ -1,11 +1,17 @@ package com.twitter.app +import com.twitter.app.lifecycle.Event._ +import com.twitter.app.lifecycle.Event +import com.twitter.app.lifecycle.Observer import com.twitter.app.LoadService.Binding -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ +import com.twitter.finagle.util.enableJvmShutdownHook import com.twitter.util._ import java.util.concurrent.ConcurrentLinkedQueue -import org.scalatest.FunSuite +import scala.jdk.CollectionConverters._ import scala.language.reflectiveCalls +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.BeforeAndAfterEach class TestApp(f: () => Unit) extends App { var reason: Option[String] = None @@ -26,21 +32,49 @@ object VeryBadApp extends App { def main(): Unit = {} } +class WeNeverClose extends App { + override def exitOnError(throwable: Throwable): Unit = { + throwable match { + case _: CloseException => + throw throwable + case _ => + exitOnError("Exception thrown in main on startup") + } + } + + private[this] val unClosable = new Closable { + override def close(deadline: Time): Future[Unit] = Future.never + } + + closeOnExit(unClosable) + + def main(): Unit = {} +} + +class WeNeverCloseButWeDoNotCare extends WeNeverClose { + override protected val suppressGracefulShutdownErrors: Boolean = true +} + trait ErrorOnExitApp extends App { - override val defaultCloseGracePeriod: Duration = 2.seconds + override val defaultCloseGracePeriod: Duration = 10.seconds override def exitOnError(throwable: Throwable): Unit = { throw throwable } } -class AppTest extends FunSuite { +class AppTest extends AnyFunSuite with BeforeAndAfterEach { + + override protected def beforeEach(): Unit = { + enableJvmShutdownHook.parse("false") + } + test("App: make sure system.exit called on exception from main") { val test1 = new TestApp(() => throw new RuntimeException("simulate main failing")) test1.main(Array()) - assert(test1.reason.contains("Exception thrown in main on startup")) + assert(test1.reason.exists(_.contains("Exception thrown in main on startup"))) } test("App: propagate underlying exception from fields in app") { @@ -77,6 +111,7 @@ class AppTest extends FunSuite { test("App: order of hooks") { val q = new ConcurrentLinkedQueue[Int] class Test1 extends App { + onExitLast(q.add(5)) onExit(q.add(4)) postmain(q.add(3)) def main(): Unit = q.add(2) @@ -84,7 +119,63 @@ class AppTest extends FunSuite { init(q.add(0)) } new Test1().main(Array.empty) - assert(q.toArray.toSeq == Seq(0, 1, 2, 3, 4)) + assert(q.asScala.toSeq == Seq(0, 1, 2, 3, 4, 5)) + } + + test("App: order of hooks with runOnExit and runOnExitLast") { + val q = new ConcurrentLinkedQueue[Int] + class Test1 extends App { + runOnExitLast(() => q.add(5)) + runOnExit(() => q.add(4)) + postmain(q.add(3)) + def main(): Unit = q.add(2) + premain(q.add(1)) + init(q.add(0)) + } + new Test1().main(Array.empty) + assert(q.asScala.toSeq == Seq(0, 1, 2, 3, 4, 5)) + } + + test("App: order of hooks via observer") { + val enterQueue = new ConcurrentLinkedQueue[Event] + val completeQueue = new ConcurrentLinkedQueue[Event] + val observer = new Observer { + override def onEntry(event: Event): Unit = enterQueue.add(event) + def onSuccess(event: Event): Unit = completeQueue.add(event) + def onFailure(event: Event, throwable: Throwable): Unit = + completeQueue.add(event) + } + + class ObservedApp extends App + val test = new ObservedApp + test.withObserver(observer) + test.main(Array.empty) + + assert( + enterQueue.asScala.toSeq == Seq( + Register, + LoadBindings, + Init, + ParseArgs, + PreMain, + Main, + PostMain, + Close, + CloseExitLast + )) + + assert( + completeQueue.asScala.toSeq == Seq( + Register, + LoadBindings, + Init, + ParseArgs, + PreMain, + Main, + PostMain, + CloseExitLast, // CloseExitLast COMPLETES within the Close phase + Close + )) } test("loadServiceBinds") { @@ -111,12 +202,8 @@ class AppTest extends FunSuite { val a = new App { def main(): Unit = () } val p = new Promise[Unit] var n1, n2 = 0 - val c1 = Closable.make { _ => - n1 += 1; p - } - val c2 = Closable.make { _ => - n2 += 1; Future.Done - } + val c1 = Closable.make { _ => n1 += 1; p } + val c2 = Closable.make { _ => n2 += 1; Future.Done } a.closeOnExitLast(c2) a.closeOnExit(c1) val f = a.close() @@ -126,7 +213,7 @@ class AppTest extends FunSuite { assert(n2 == 0) assert(!f.isDefined) - p.setDone + p.setDone() assert(n2 == 1) assert(f.isDefined) } @@ -140,12 +227,8 @@ class AppTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => var n1, n2 = 0 - val c1 = Closable.make { _ => - n1 += 1; Future.never - } - val c2 = Closable.make { _ => - n2 += 1; Future.never - } + val c1 = Closable.make { _ => n1 += 1; Future.never } + val c2 = Closable.make { _ => n2 += 1; Future.never } a.closeOnExitLast(c2) a.closeOnExit(c1) val f = a.close(Time.now + 1.second) @@ -175,18 +258,10 @@ class AppTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => var n1, n2, n3, n4 = 0 - val c1 = Closable.make { _ => - n1 += 1; Future.never - } // Exit - val c2 = Closable.make { _ => - n2 += 1; Future.Done - } // ExitLast - val c3 = Closable.make { _ => - n3 += 1; Future.Done - } // Late Exit - val c4 = Closable.make { _ => - n4 += 1; Future.Done - } // Late ExitLast + val c1 = Closable.make { _ => n1 += 1; Future.never } // Exit + val c2 = Closable.make { _ => n2 += 1; Future.Done } // ExitLast + val c3 = Closable.make { _ => n3 += 1; Future.Done } // Late Exit + val c4 = Closable.make { _ => n4 += 1; Future.Done } // Late ExitLast a.closeOnExitLast(c2) a.closeOnExit(c1) val f = a.close(Time.now + 1.second) @@ -227,18 +302,10 @@ class AppTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => var n1, n2, n3, n4 = 0 - val c1 = Closable.make { _ => - n1 += 1; Future.never - } // Exit - val c2 = Closable.make { _ => - n2 += 1; Future.never - } // ExitLast - val c3 = Closable.make { _ => - n3 += 1; Future.never - } // Late Exit - val c4 = Closable.make { _ => - n4 += 1; Future.never - } // Late ExitLast + val c1 = Closable.make { _ => n1 += 1; Future.never } // Exit + val c2 = Closable.make { _ => n2 += 1; Future.never } // ExitLast + val c3 = Closable.make { _ => n3 += 1; Future.never } // Late Exit + val c4 = Closable.make { _ => n4 += 1; Future.never } // Late ExitLast a.closeOnExitLast(c2) a.closeOnExit(c1) val f = a.close(Time.now + 1.second) @@ -306,17 +373,13 @@ class AppTest extends FunSuite { test("App: exit functions properly capture non-fatal exceptions") { val app = new ErrorOnExitApp { def main(): Unit = { - onExit{ + onExit { throw new Exception("FORCED ON EXIT") } - closeOnExit(Closable.make { _ => - throw new Exception("FORCED CLOSE ON EXIT") - }) + closeOnExit(Closable.make { _ => throw new Exception("FORCED CLOSE ON EXIT") }) - closeOnExitLast(Closable.make { _ => - throw new Exception("FORCED CLOSE ON EXIT LAST") - }) + closeOnExitLast(Closable.make { _ => throw new Exception("FORCED CLOSE ON EXIT LAST") }) } } @@ -331,9 +394,7 @@ class AppTest extends FunSuite { // first fatal (InterruptedException) kills the app during close val app = new ErrorOnExitApp { def main(): Unit = { - closeOnExit(Closable.make { _ => - throw new InterruptedException("FORCED CLOSE ON EXIT") - }) + closeOnExit(Closable.make { _ => throw new InterruptedException("FORCED CLOSE ON EXIT") }) } } @@ -345,13 +406,11 @@ class AppTest extends FunSuite { test("App: exit functions properly capture mix of non-fatal and fatal exceptions") { val app = new ErrorOnExitApp { def main(): Unit = { - onExit{ + onExit { throw new Exception("FORCED ON EXIT") } - closeOnExit(Closable.make { _ => - throw new Exception("FORCED CLOSE ON EXIT") - }) + closeOnExit(Closable.make { _ => throw new Exception("FORCED CLOSE ON EXIT") }) closeOnExitLast(Closable.make { _ => throw new InterruptedException("FORCED CLOSE ON EXIT LAST") @@ -376,4 +435,38 @@ class AppTest extends FunSuite { app.main(Array.empty) } + + test("never closing exits with exception") { + val t = new MockTimer + val app = new WeNeverClose { + override lazy val shutdownTimer: Timer = t + } + + Time.withTimeAt(Time.epoch) { _ => + val exc = intercept[CloseException] { + app.main(Array.empty) + } + + // the CloseException should be suppressing a com.twitter.util.TimeoutException + assert(exc.getSuppressed.length == 1) + assert(exc.getSuppressed.head.getClass == classOf[com.twitter.util.TimeoutException]) + } + } + + test("never closing exits with a timeout exception") { + val t = new MockTimer + val app = new WeNeverCloseButWeDoNotCare { + override lazy val shutdownTimer: Timer = t + } + + Time.withTimeAt(Time.epoch) { _ => + val exc = intercept[CloseException] { + app.main(Array.empty) + } + + // the CloseException should be suppressing a com.twitter.util.TimeoutException + assert(exc.getSuppressed.length == 1) + assert(exc.getSuppressed.head.getClass == classOf[com.twitter.util.TimeoutException]) + } + } } diff --git a/util-app/src/test/scala/com/twitter/app/ClassPathTest.scala b/util-app/src/test/scala/com/twitter/app/ClassPathTest.scala index 93984fddb2..d2752fbb85 100644 --- a/util-app/src/test/scala/com/twitter/app/ClassPathTest.scala +++ b/util-app/src/test/scala/com/twitter/app/ClassPathTest.scala @@ -2,10 +2,10 @@ package com.twitter.app import java.net.URLClassLoader import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -class ClassPathTest extends FunSuite with MockitoSugar { +class ClassPathTest extends AnyFunSuite with MockitoSugar { test("Null URL[] URLClassloader") { val classLoader = mock[URLClassLoader] diff --git a/util-app/src/test/scala/com/twitter/app/ColonizeTestApp.scala b/util-app/src/test/scala/com/twitter/app/ColonizeTestApp.scala index efb2c270bc..295788062a 100644 --- a/util-app/src/test/scala/com/twitter/app/ColonizeTestApp.scala +++ b/util-app/src/test/scala/com/twitter/app/ColonizeTestApp.scala @@ -7,8 +7,8 @@ object SolarSystemPlanets { val orderFromSun: Int, val name: String, val mass: Kilogram, - val radius: Meter - ) extends Ordered[Planet] { + val radius: Meter) + extends Ordered[Planet] { def compare(that: Planet): Int = this.orderFromSun - that.orderFromSun @@ -50,7 +50,7 @@ object SolarSystemPlanets { type Kilogram = Double type Meter = Double - private val G = 6.67300E-11 // universal gravitational constant (m3 kg-1 s-2) + private val G = 6.67300e-11 // universal gravitational constant (m3 kg-1 s-2) } class ColonizeTestApp extends App { diff --git a/util-app/src/test/scala/com/twitter/app/FlagTest.scala b/util-app/src/test/scala/com/twitter/app/FlagTest.scala index 263aaeb898..940c73ff3a 100644 --- a/util-app/src/test/scala/com/twitter/app/FlagTest.scala +++ b/util-app/src/test/scala/com/twitter/app/FlagTest.scala @@ -1,17 +1,18 @@ package com.twitter.app import com.twitter.app.SolarSystemPlanets._ -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, Awaitable} -import org.scalatest.FunSuite import scala.collection.mutable +import org.scalatest.funsuite.AnyFunSuite -class FlagTest extends FunSuite { +class FlagTest extends AnyFunSuite { class Ctx(failFastUntilParsed: Boolean = false) { val flag = new Flags("test", includeGlobal = false, failFastUntilParsed) val fooFlag = flag("foo", 123, "The foo value") val barFlag = flag("bar", "okay", "The bar value") + val bazFlag = flag[String]("baz", "The baz value") } test("Flag: defaults") { @@ -21,6 +22,11 @@ class FlagTest extends FunSuite { flag.finishParsing() assert(fooFlag() == 123) assert(barFlag() == "okay") + intercept[IllegalArgumentException](bazFlag()) + + assert(fooFlag.getDefault == Some(123)) + assert(barFlag.getDefault == Some("okay")) + assert(bazFlag.getDefault.isEmpty) } test("Flag: override a flag") { @@ -86,6 +92,38 @@ class FlagTest extends FunSuite { assert(res2 == "false") } + test("Flag: letParse") { + def current: Boolean = MyGlobalBooleanFlag() + + // track the order the blocks execute and that they only execute once + var buf = mutable.Buffer[Int]() + + // make sure they stack properly + assert(!current) + MyGlobalBooleanFlag.letParse("true") { + buf += 1 + assert(current) + MyGlobalBooleanFlag.letParse("false") { + buf += 2 + assert(!current) + } + buf += 3 + assert(current) + } + assert(!current) + + assert(buf == Seq(1, 2, 3)) + } + + test("Flag: letParse return values reflect bindings") { + def current: Boolean = MyGlobalBooleanFlag() + + val res1 = MyGlobalBooleanFlag.letParse("true") { current.toString } + assert(res1 == "true") + val res2 = MyGlobalBooleanFlag.letParse("false") { current.toString } + assert(res2 == "false") + } + test("Flag: letClear") { // track the order the blocks execute and that they only execute once var buf = mutable.Buffer[Int]() @@ -143,6 +181,15 @@ class FlagTest extends FunSuite { assert(noDefaultAndSupplied.get.contains(3)) } + test("Flag.getUnparsed") { + val ctx = new GetCtx() + import ctx._ + + assert(withDefault.getUnparsed.isEmpty) + assert(noDefault.getUnparsed.isEmpty) + assert(noDefaultAndSupplied.getUnparsed.contains("3")) + } + test("Flag.getWithDefault") { val ctx = new GetCtx() import ctx._ @@ -152,6 +199,15 @@ class FlagTest extends FunSuite { assert(noDefaultAndSupplied.getWithDefault.contains(3)) } + test("Flag.getWithDefaultUnparsed") { + val ctx = new GetCtx() + import ctx._ + + assert(withDefault.getWithDefaultUnparsed.contains("1")) + assert(noDefault.getWithDefaultUnparsed.isEmpty) + assert(noDefaultAndSupplied.getWithDefaultUnparsed.contains("3")) + } + test("Flag.usage - new flag value") { val marsColony = new ColonizeTestApp() try { diff --git a/util-app/src/test/scala/com/twitter/app/FlaggableTest.scala b/util-app/src/test/scala/com/twitter/app/FlaggableTest.scala index bf34c7b067..91bddb83bf 100644 --- a/util-app/src/test/scala/com/twitter/app/FlaggableTest.scala +++ b/util-app/src/test/scala/com/twitter/app/FlaggableTest.scala @@ -1,8 +1,11 @@ package com.twitter.app -import org.scalatest.FunSuite +import com.twitter.util.Duration +import com.twitter.util.Time +import java.time.LocalTime +import org.scalatest.funsuite.AnyFunSuite -class FlaggableTest extends FunSuite { +class FlaggableTest extends AnyFunSuite { test("Flaggable: parse booleans") { assert(Flaggable.ofBoolean.parse("true")) @@ -16,6 +19,16 @@ class FlaggableTest extends FunSuite { assert(Flaggable.ofString.parse("blah") == "blah") } + test("Flaggagle: parse dates and times") { + val t = Time.fromSeconds(123) + assert(Flaggable.ofTime.parse(t.toString) == t) + + assert(Flaggable.ofDuration.parse("1.second") == Duration.fromSeconds(1)) + + val tt = LocalTime.ofSecondOfDay(123) + assert(Flaggable.ofJavaLocalTime.parse(tt.toString) == tt) + } + test("Flaggable: parse/show inet addresses") { // we aren't binding to this port, so any one will do. val port = 1589 @@ -23,10 +36,10 @@ class FlaggableTest extends FunSuite { assert(local.getAddress.isAnyLocalAddress) assert(local.getPort == port) - val ip = "141.211.133.111" + val ip = "127.0.0.1" val remote = Flaggable.ofInetSocketAddress.parse(s"$ip:$port") - assert(remote.getHostName == remote.getHostName) + assert(remote.getHostName == "localhost") assert(remote.getPort == port) assert(Flaggable.ofInetSocketAddress.show(local) == s":$port") @@ -35,28 +48,41 @@ class FlaggableTest extends FunSuite { test("Flaggable: parse seqs") { assert(Flaggable.ofSeq[Int].parse("1,2,3,4") == Seq(1, 2, 3, 4)) + assert(Flaggable.ofSeq[Int].parse("") == Seq.empty[Int]) + assert(Flaggable.ofSeq[Boolean].parse("true,false,true") == Seq(true, false, true)) + } + + test("Flaggable: parse sets") { + assert(Flaggable.ofSet[Int].parse("1,2,3,4") == Set(1, 2, 3, 4)) + assert(Flaggable.ofSet[Int].parse("") == Set.empty[Int]) + assert(Flaggable.ofSeq[Boolean].parse("true,false,true") == Seq(true, false, true)) } test("Flaggable: parse maps") { assert(Flaggable.ofMap[Int, Int].parse("1=2,3=4") == Map(1 -> 2, 3 -> 4)) + assert( + Flaggable + .ofMap[String, Boolean].parse("foo=true,bar=false") == Map("foo" -> true, "bar" -> false)) } test("Flaggable: parse maps with comma-separated values") { assert( - Flaggable.ofMap[String, Seq[Int]].parse("a=1,2,3,3,b=4,5") == + Flaggable.ofMap[String, Seq[Int]].parse("a=1,2,3,3,b=4,5") == // 1000191039 Map("a" -> Seq(1, 2, 3, 3), "b" -> Seq(4, 5)) ) } - test("Flaggable: parse maps of sets with comma-separated values") { - assert( - Flaggable.ofMap[String, Set[Int]].parse("a=1,2,3,3,b=4,5") == - Map("a" -> Set(1, 2, 3), "b" -> Set(4, 5)) - ) - } - test("Flaggable: parse tuples") { assert(Flaggable.ofTuple[Int, String].parse("1,hello") == ((1, "hello"))) + assert(Flaggable.ofTuple[String, Boolean].parse("foo,true") == (("foo", true))) intercept[IllegalArgumentException] { Flaggable.ofTuple[Int, String].parse("1") } } + + test("Flaggable: expose types and kinds") { + assert(Flaggable.ofFile.isInstanceOf[Flaggable.Typed[_]]) + assert(Flaggable.ofJavaBoolean.isInstanceOf[Flaggable.Typed[_]]) + + assert(Flaggable.ofSeq[Int].isInstanceOf[Flaggable.Generic[_]]) + assert(Flaggable.ofJavaSet[Int].isInstanceOf[Flaggable.Generic[_]]) + } } diff --git a/util-app/src/test/scala/com/twitter/app/FlagsTest.scala b/util-app/src/test/scala/com/twitter/app/FlagsTest.scala index 2609ad28c8..cdb5ca95fe 100644 --- a/util-app/src/test/scala/com/twitter/app/FlagsTest.scala +++ b/util-app/src/test/scala/com/twitter/app/FlagsTest.scala @@ -1,9 +1,11 @@ package com.twitter.app -import com.twitter.util.registry.{Entry, GlobalRegistry, SimpleRegistry} -import org.scalatest.FunSuite +import com.twitter.util.registry.Entry +import com.twitter.util.registry.GlobalRegistry +import com.twitter.util.registry.SimpleRegistry +import org.scalatest.funsuite.AnyFunSuite -class FlagsTest extends FunSuite { +class FlagsTest extends AnyFunSuite { class Ctx(failFastUntilParsed: Boolean = false) { val flag = new Flags("test", includeGlobal = false, failFastUntilParsed) @@ -165,6 +167,18 @@ class FlagsTest extends FunSuite { ) } + test("Flag: multiple parse errors roll-up") { + val ctx = new Ctx + import ctx._ + flag.parseArgs(Array("-undefined", "-foo", "blah")) match { + case Flags.Error(reason) => + assert(reason.contains("Error parsing flag \"undefined\": " + Flags.FlagUndefinedMessage)) + assert(reason.contains("Error parsing flag \"foo\": For input string: \"blah\"")) + assert(reason.contains("usage:")) + case other => fail(s"expected a flag error, but received $other") + } + } + test("formatFlagValues") { val flagWithGlobal = new Flags("my", includeGlobal = true) @@ -254,12 +268,19 @@ class FlagsTest extends FunSuite { new Flag[Int]("something.id", "bar", Left(() => 3), failFastUntilParsed = false) assert(somethingIdFlag() == 3) // this logs a message in SEVERE and we get the default value - val testFlags = new Flags("failFastTest", includeGlobal = false, failFastUntilParsed = true) // added flags will inherit this value + val testFlags = + new Flags( + "failFastTest", + includeGlobal = false, + failFastUntilParsed = true + ) // added flags will inherit this value testFlags.add(somethingIdFlag) intercept[IllegalStateException] { testFlags .getAll() - .foreach(flag => if (flag.name == "something.id") flag.apply()) // this should fail, since all added flags should now fail fast. + .foreach(flag => + if (flag.name == "something.id") + flag.apply()) // this should fail, since all added flags should now fail fast. } // now parse and apply diff --git a/util-app/src/test/scala/com/twitter/app/GlobalFlagTest.scala b/util-app/src/test/scala/com/twitter/app/GlobalFlagTest.scala index 101a3f78f9..f17cb34d8e 100644 --- a/util-app/src/test/scala/com/twitter/app/GlobalFlagTest.scala +++ b/util-app/src/test/scala/com/twitter/app/GlobalFlagTest.scala @@ -1,12 +1,13 @@ package com.twitter.app -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite object MyGlobalFlag extends GlobalFlag[String]("a test flag", "a global test flag") object MyGlobalFlagNoDefault extends GlobalFlag[Int]("a global test flag with no default") object MyGlobalBooleanFlag extends GlobalFlag[Boolean](false, "a boolean flag") -class GlobalFlagTest extends FunSuite { +class GlobalFlagTest extends AnyFunSuite { + val flagSet = Set(MyGlobalFlag, MyGlobalBooleanFlag, MyGlobalFlagNoDefault) test("GlobalFlag.get") { assert(MyGlobalBooleanFlag.get.isEmpty) @@ -63,4 +64,20 @@ class GlobalFlagTest extends FunSuite { } } + test("GlobalFlag.getAll") { + val mockClassLoader = new MockClassLoader(getClass.getClassLoader) + val flags = GlobalFlag.getAll(mockClassLoader) + assert(flags.toSet == flagSet) + } + + private class MockClassLoader(realClassLoader: ClassLoader) extends ClassLoader(realClassLoader) { + private val isValidClassName = (className: String) => + flagSet + .map(_.getClass.getName) + .contains(className) + + override def loadClass(name: String): Class[_] = + if (isValidClassName(name)) realClassLoader.loadClass(name) else null + } + } diff --git a/util-app/src/test/scala/com/twitter/app/LoadServiceTest.scala b/util-app/src/test/scala/com/twitter/app/LoadServiceTest.scala index 5e83545800..d4aa02f791 100644 --- a/util-app/src/test/scala/com/twitter/app/LoadServiceTest.scala +++ b/util-app/src/test/scala/com/twitter/app/LoadServiceTest.scala @@ -1,7 +1,7 @@ package com.twitter.app import com.twitter.app.LoadService.Binding -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.finagle.util.loadServiceDenied import com.twitter.util.registry.{Entry, GlobalRegistry, SimpleRegistry} import com.twitter.util.{Await, Future, FuturePool} @@ -9,10 +9,10 @@ import java.io.{File, InputStream} import java.net.URL import java.util import java.util.concurrent.{Callable, CountDownLatch, ExecutorService, Executors} -import java.util.{Random, concurrent} -import org.scalatest.FunSuite -import org.scalatest.mockito.MockitoSugar -import scala.collection.mutable +import java.util.concurrent.{Future => JFuture} +import java.util.Random +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite // These traits correspond to files in: // util-app/src/test/resources/META-INF/services @@ -51,7 +51,7 @@ class LoadServiceDeadlockImpl extends LoadServiceDeadlock { UsesLoadService.lsd } -class LoadServiceTest extends FunSuite with MockitoSugar { +class LoadServiceTest extends AnyFunSuite with MockitoSugar { test("LoadService should apply[T] and return a set of instances of T") { assert(LoadService[LoadServiceRandomInterface]().nonEmpty) @@ -72,7 +72,8 @@ class LoadServiceTest extends FunSuite with MockitoSugar { } } - test("LoadService should only load 1 instance of T, even when there's multiple occurrences of T") { + test( + "LoadService should only load 1 instance of T, even when there's multiple occurrences of T") { val randomIfaces = LoadService[LoadServiceRandomInterface]() assert(randomIfaces.size == 1) } @@ -99,7 +100,7 @@ class LoadServiceTest extends FunSuite with MockitoSugar { test("LoadService shouldn't fail on un-readable dir") { val loader = mock[ClassLoader] - val buf = mutable.Buffer.empty[ClassPath.LoadServiceInfo] + val buf = Vector.newBuilder[ClassPath.LoadServiceInfo] val rand = new Random() val f = File.createTempFile("tmp", "__utilapp_loadservice" + rand.nextInt(10000)) @@ -108,14 +109,14 @@ class LoadServiceTest extends FunSuite with MockitoSugar { f.setReadable(false) new LoadServiceClassPath().browseUri(f.toURI, loader, buf) - assert(buf.isEmpty) + assert(buf.result().isEmpty) f.delete() } } test("LoadService shouldn't fail on un-readable sub-dir") { val loader = mock[ClassLoader] - val buf = mutable.Buffer.empty[ClassPath.LoadServiceInfo] + val buf = Vector.newBuilder[ClassPath.LoadServiceInfo] val rand = new Random() val f = File.createTempFile("tmp", "__utilapp_loadservice" + rand.nextInt(10000)) @@ -127,7 +128,7 @@ class LoadServiceTest extends FunSuite with MockitoSugar { subDir.setReadable(false) new LoadServiceClassPath().browseUri(f.toURI, loader, buf) - assert(buf.isEmpty) + assert(buf.result().isEmpty) subDir.delete() f.delete() @@ -140,7 +141,7 @@ class LoadServiceTest extends FunSuite with MockitoSugar { // Run LoadService in a different thread from the custom classloader val clazz: Class[_] = loader.loadClass("com.twitter.app.LoadServiceCallable") val executor: ExecutorService = Executors.newSingleThreadExecutor() - val future: concurrent.Future[Seq[Any]] = + val future: JFuture[Seq[Any]] = executor.submit(clazz.newInstance().asInstanceOf[Callable[Seq[Any]]]) // Get the result @@ -165,7 +166,7 @@ class LoadServiceTest extends FunSuite with MockitoSugar { val jos = new JarOutputStream(new FileOutputStream(jarFile), manifest) jos.close() val loader = mock[ClassLoader] - val buf = mutable.Buffer.empty[ClassPath.LoadServiceInfo] + val buf = Vector.newBuilder[ClassPath.LoadServiceInfo] new LoadServiceClassPath().browseUri(jarFile.toURI, loader, buf) } finally { jarFile.delete @@ -187,7 +188,7 @@ class LoadServiceTest extends FunSuite with MockitoSugar { attributes.put(CLASS_PATH, jar1.getName) new JarOutputStream(new FileOutputStream(jar2), manifest).close() val loader = mock[ClassLoader] - val buf = mutable.Buffer.empty[ClassPath.LoadServiceInfo] + val buf = Vector.newBuilder[ClassPath.LoadServiceInfo] new LoadServiceClassPath().browseUri(jar1.toURI, loader, buf) } finally { jar1.delete @@ -262,7 +263,6 @@ class LoadServiceTest extends FunSuite with MockitoSugar { assert(lsds0.size == 1) assert(lsds0.head.getClass == classOf[LoadServiceDeadlockImpl]) } - } class LoadServiceCallable extends Callable[Seq[Any]] { diff --git a/util-app/src/test/scala/com/twitter/app/command/CommandTest.scala b/util-app/src/test/scala/com/twitter/app/command/CommandTest.scala new file mode 100644 index 0000000000..b07851a15e --- /dev/null +++ b/util-app/src/test/scala/com/twitter/app/command/CommandTest.scala @@ -0,0 +1,251 @@ +package com.twitter.app.command + +import com.twitter.conversions.DurationOps._ +import com.twitter.io.{Buf, Reader} +import com.twitter.util.{Await, Awaitable, Promise} +import java.io.{File, InputStream, OutputStream, PipedInputStream, PipedOutputStream} +import org.scalatest.BeforeAndAfterEach +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers + +class CommandTest extends AnyFunSuite with BeforeAndAfterEach with Matchers { + private def await[A](awaitable: Awaitable[A]): A = Await.result(awaitable, 5.seconds) + + private class FakeCommandExecutorImpl extends CommandExecutor { + private var retValue: Promise[Int] = _ + private var command: Seq[String] = _ + private var workingDirectory: Option[File] = _ + private var errorOutputStream: PipedOutputStream = _ + private var errorInputStream: PipedInputStream = _ + private var stdOutputStream: PipedOutputStream = _ + private var stdInputStream: PipedInputStream = _ + + override def apply( + cmd: Seq[String], + directory: Option[File], + extraEnv: Map[String, String] + ): Process = { + retValue = new Promise[Int]() + command = cmd + workingDirectory = directory + errorOutputStream = new PipedOutputStream() + errorInputStream = new PipedInputStream(errorOutputStream) + stdOutputStream = new PipedOutputStream() + stdInputStream = new PipedInputStream(stdOutputStream) + + new Process { + override def getOutputStream: OutputStream = ??? + + override def getInputStream: InputStream = stdInputStream + + override def getErrorStream: InputStream = errorInputStream + + override def waitFor(): Int = { + Await.result(retValue, 1.second) + } + + override def exitValue(): Int = ??? + + override def destroy(): Unit = ??? + } + } + + private def checkInitialized(): Unit = { + if (retValue == null || errorInputStream == null || stdInputStream == null || command == null) { + throw new IllegalStateException( + "Can't use fake command executor until the command has been kicked off with apply()") + } + } + + private def writeToStream(buf: Buf, stream: OutputStream): Unit = { + stream.write(Buf.ByteArray.Shared.extract(buf)) + } + + def writeToStdOut(line: Buf): Unit = { + checkInitialized() + writeToStream(line, stdOutputStream) + } + + def writeToStdErr(line: Buf): Unit = { + checkInitialized() + writeToStream(line, errorOutputStream) + } + + def setReturn(value: Int): Unit = { + checkInitialized() + retValue.setValue(value) + errorOutputStream.flush() + stdOutputStream.flush() + errorOutputStream.close() + stdOutputStream.close() + } + + def reset(): Unit = { + command = null + retValue = null + errorOutputStream = null + errorInputStream = null + stdOutputStream = null + stdInputStream = null + } + } + + private val fakeCommandExecutor = new FakeCommandExecutorImpl + private val exposedCommandExecutor: CommandExecutor = fakeCommandExecutor + private val command = new Command(exposedCommandExecutor) + + private def fakeCommandReturnsLines(lines: Buf*): Unit = { + lines.foreach(fakeCommandExecutor.writeToStdOut) + } + + private def fakeCommandReturnsValue(retVal: Int): Unit = { + fakeCommandExecutor.setReturn(retVal) + } + + private def fakeCommandReturnsStdErr(stdErrLines: Buf*): Unit = { + stdErrLines.foreach(fakeCommandExecutor.writeToStdErr) + } + + private def parseUtf8Buf(buf: Buf): String = buf match { + case Buf.Utf8(str) => str + } + + override protected def beforeEach(): Unit = { + fakeCommandExecutor.reset() + } + + test("When command succeeds return reader of lines with stdout and stderr") { + val result = command.run(Seq("command")) + val expectedStdOut = Seq("line 1\n", "line 2\n").map(Buf.Utf8(_)) + val expectedStdErr = Seq("error 1\n", "error 2\n").map(Buf.Utf8(_)) + fakeCommandReturnsLines(expectedStdOut: _*) + fakeCommandReturnsStdErr(expectedStdErr: _*) + fakeCommandReturnsValue(0) + val actualStdOut = await(Reader.readAllItems(result.stdout)) + actualStdOut.map(parseUtf8Buf) shouldBe Seq("line 1", "line 2") + + val actualStdErr = await(Reader.readAllItems(result.stderr)) + actualStdErr.map(parseUtf8Buf) shouldBe Seq("error 1", "error 2") + } + + test("When command fails at end return reader of lines") { + val result = command.run(Seq("command")) + val expectedStdOut = Seq("line 1\n", "line 2\n").map(Buf.Utf8(_)) + val expectedStdErr = Seq("error 1\n", "error 2\n").map(Buf.Utf8(_)) + fakeCommandReturnsLines(expectedStdOut: _*) + fakeCommandReturnsStdErr(expectedStdErr: _*) + fakeCommandReturnsValue(-1) + val stdOutLine1 = await(result.stdout.read()) + val stdOutLine2 = await(result.stdout.read()) + stdOutLine1.map(parseUtf8Buf) shouldBe Some("line 1") + stdOutLine2.map(parseUtf8Buf) shouldBe Some("line 2") + + val stdErrLine1 = await(result.stderr.read()) + stdErrLine1.map(parseUtf8Buf) shouldBe Some("error 1") + + val ex = the[NonZeroReturnCode] thrownBy { + await(result.stdout.read()) + } + + // We don't get the first line because it was already read + await(Reader.readAllItems(ex.stdErr)).map(parseUtf8Buf) shouldBe Seq("error 2") + ex.code shouldBe -1 + } + + test("splitOnNewlinees prpoerly a simple single line with no newline") { + assertSplitOnNewlines(Seq("line 1"), Seq.empty, Some("line 1")) + } + + test("splitOnNewlinees prpoerly splits up a simple single line with a newline") { + assertSplitOnNewlines(Seq("line 1\n"), Seq("line 1"), None) + } + + test("splitOnNewlinees prpoerly splits up a single line split into 3 without a newline") { + assertSplitOnNewlines(Seq("li", "ne", " 1"), Seq.empty, Some("line 1")) + } + + test("splitOnNewlinees prpoerly splits up a single line split into 3 with a newline") { + assertSplitOnNewlines(Seq("li", "ne", " 1\n"), Seq("line 1"), None) + } + + test("splitOnNewlinees splits up a simple single line with a newline in a different segment") { + assertSplitOnNewlines(Seq("line 1", "\n"), Seq("line 1"), None) + } + + test("splitOnNewlinees splits up an empty input") { + assertSplitOnNewlines(Seq.empty, Seq.empty, None) + } + + test("splitOnNewlinees splits up an empty string input") { + assertSplitOnNewlines(Seq(""), Seq.empty, None) + } + + test("splitOnNewlinees splits up just a newline") { + assertSplitOnNewlines(Seq("\n"), Seq(""), None) + } + + test("splitOnNewlinees splits up a multiple lines in one segment with no trailing newline") { + assertSplitOnNewlines(Seq("line 1\nline 2"), Seq("line 1"), Some("line 2")) + } + + test("splitOnNewlinees splits up a multiple lines in one segment with a trailing newline") { + assertSplitOnNewlines(Seq("line 1\nline 2\n"), Seq("line 1", "line 2"), None) + } + + test( + "splitOnNewlinees splits up a multiple lines in multiple segments with no trailing newline") { + assertSplitOnNewlines( + Seq("li", "ne 1\nli", "ne 2\nfinal text"), + Seq("line 1", "line 2"), + Some("final text")) + } + + test("splitOnNewlinees splits up a multiple lines in multiple segments with a trailing newline") { + assertSplitOnNewlines( + Seq("li", "ne 1\nli", "ne 2\nfinal text\n"), + Seq("line 1", "line 2", "final text"), + None) + } + + /** + * A helper function to test many cases around splitOnNewlines. + * + * @param input The strings used to generate the input [[Reader]] + * @param expectedMainOutput What we expect to read from the [[Reader]] before the input [[Reader]] is + * finished + * @param expectedLastLine What we expect to read from the [[Reader]] once the input [[Reader]] is + * finished + */ + private def assertSplitOnNewlines( + input: Seq[String], + expectedMainOutput: Seq[String], + expectedLastLine: Option[String] + ): Unit = { + val finishingPromise = new Promise[Buf]() + val inputReader = Reader.concat( + Seq(Reader.fromSeq(input.map(Buf.Utf8(_))), Reader.fromFuture(finishingPromise))) + + val splitReader = CommandOutput.splitOnNewLines(inputReader) + + // Check main output + expectedMainOutput.foreach { expectedLine => + await(splitReader.read()).map(parseUtf8Buf) shouldBe Some(expectedLine) + } + + // The last line isn't available yet because the reader is not yet done + val finalTextFut = splitReader.read() + assert(!finalTextFut.isDefined) + + finishingPromise.setValue(Buf.Empty) + + // Now we should have a line read, if it's available + await(finalTextFut).map(parseUtf8Buf) shouldBe expectedLastLine + + // If there was content then read again to ensure we are done + if (expectedLastLine.isDefined) { + await(splitReader.read()) shouldBe None + } + + } + +} diff --git a/util-app/src/test/scala/com/twitter/app/command/RealCommandTest.scala b/util-app/src/test/scala/com/twitter/app/command/RealCommandTest.scala new file mode 100644 index 0000000000..72fd8ea805 --- /dev/null +++ b/util-app/src/test/scala/com/twitter/app/command/RealCommandTest.scala @@ -0,0 +1,89 @@ +package com.twitter.app.command + +import com.twitter.app.command.RealCommandTest.ScriptPath +import com.twitter.conversions.DurationOps._ +import com.twitter.io.{Buf, Reader, TempFile} +import com.twitter.util.{Await, Awaitable, Future} +import java.nio.file.attribute.PosixFilePermission +import java.nio.file.Files +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import scala.jdk.CollectionConverters._ + +object RealCommandTest { + // Bazel provides the internal dependency as a JAR, and consequently, + // we need to copy the script into a temp directory to use it. + // http://go/bazel-compatibility/no-such-file + private val ScriptPath: String = { + val temp = TempFile.fromResourcePath(this.getClass, "/command-test.sh") + Files.setPosixFilePermissions( + temp.toPath, + Set( + PosixFilePermission.OWNER_READ, + PosixFilePermission.OWNER_EXECUTE, + PosixFilePermission.GROUP_READ, + PosixFilePermission.GROUP_EXECUTE, + PosixFilePermission.OTHERS_READ, + PosixFilePermission.OTHERS_EXECUTE + ).asJava + ) + temp.toString + } +} + +class RealCommandTest extends AnyFunSuite with Matchers { + private def await[A](awaitable: Awaitable[A]): A = Await.result(awaitable, 5.seconds) + + private def parseUtf8Buf(buf: Buf): String = { + val Buf.Utf8(str) = buf + str + } + + test("Executes a script and gets output") { + val output = Command.run( + Seq(ScriptPath, "10", "0.1", "0"), + None, + Map( + "EXTRA_ENV" -> "test value" + )) + + val firstLine = await(output.stdout.read()) + val firstError = await(output.stderr.read()) + firstLine.map(parseUtf8Buf) shouldBe Some("Stdout # 1 test value") + firstError.map(parseUtf8Buf) shouldBe Some("Stderr # 1 test value") + + // rest of lines/errors + val restOut = await(Reader.readAllItems(output.stdout)) + val restErr = await(Reader.readAllItems(output.stderr)) + restOut.map(parseUtf8Buf) shouldBe (2 to 10).map(rep => s"Stdout # $rep test value") + restErr.map(parseUtf8Buf) shouldBe (2 to 10).map(rep => s"Stderr # $rep test value") + } + + test("Executes a script and gets failure") { + val output = + Command.run( + Seq(ScriptPath, "10", "0.1", "1"), + None, + Map( + "EXTRA_ENV" -> "test value" + )) + // read 10 lines + val tenLines = await(Future.traverseSequentially(1 to 10) { _ => + output.stdout.read() + }) + + // Read a line from stderr + val stdErrLine = await(output.stderr.read()) + stdErrLine.map(parseUtf8Buf) shouldBe Some("Stderr # 1 test value") + + tenLines.flatten.map(parseUtf8Buf) shouldBe (1 to 10).map(rep => s"Stdout # $rep test value") + // read last exception + val ex = the[NonZeroReturnCode] thrownBy { await(output.stdout.read()) } + ex.code shouldBe 1 + + // Note that stdErr doesn't have the first line, because it was already read + await(Reader.readAllItems(ex.stdErr)).map(parseUtf8Buf) shouldBe (2 to 10).map { rep => + s"Stderr # $rep test value" + } + } +} diff --git a/util-benchmark/BUILD b/util-benchmark/BUILD deleted file mode 100644 index a3e0646ab6..0000000000 --- a/util-benchmark/BUILD +++ /dev/null @@ -1,5 +0,0 @@ -target( - dependencies = [ - "util/util-benchmark/src/main/scala", - ], -) diff --git a/util-benchmark/OWNERS b/util-benchmark/OWNERS deleted file mode 100644 index c313f6c1d5..0000000000 --- a/util-benchmark/OWNERS +++ /dev/null @@ -1,5 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta diff --git a/util-benchmark/PROJECT b/util-benchmark/PROJECT new file mode 100644 index 0000000000..009bd1c5fb --- /dev/null +++ b/util-benchmark/PROJECT @@ -0,0 +1,5 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan diff --git a/util-benchmark/src/main/java/com/twitter/json/BUILD.bazel b/util-benchmark/src/main/java/com/twitter/json/BUILD.bazel new file mode 100644 index 0000000000..de1e74b94f --- /dev/null +++ b/util-benchmark/src/main/java/com/twitter/json/BUILD.bazel @@ -0,0 +1,5 @@ +java_library( + sources = [ + "**/*.java", + ], +) diff --git a/util-benchmark/src/main/java/com/twitter/json/TestDemographic.java b/util-benchmark/src/main/java/com/twitter/json/TestDemographic.java new file mode 100644 index 0000000000..63808a4dc6 --- /dev/null +++ b/util-benchmark/src/main/java/com/twitter/json/TestDemographic.java @@ -0,0 +1,6 @@ +package com.twitter.json; + +public enum TestDemographic { + gender, + age; +} diff --git a/util-benchmark/src/main/java/com/twitter/json/TestFormat.java b/util-benchmark/src/main/java/com/twitter/json/TestFormat.java new file mode 100644 index 0000000000..ca8387f78f --- /dev/null +++ b/util-benchmark/src/main/java/com/twitter/json/TestFormat.java @@ -0,0 +1,7 @@ +package com.twitter.json; + +public enum TestFormat { + json, + csv, + tsv; +} diff --git a/util-benchmark/src/main/scala/BUILD b/util-benchmark/src/main/scala/BUILD deleted file mode 100644 index 88ff37f0f1..0000000000 --- a/util-benchmark/src/main/scala/BUILD +++ /dev/null @@ -1,20 +0,0 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", - "3rdparty/jvm/org/openjdk/jmh:jmh-core", - "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-core/src/main/scala", - "util/util-hashing/src/main/scala", - "util/util-stats/src/main/scala", - ], -) - -jvm_binary( - name = "jmh", - main = "org.openjdk.jmh.Main", - dependencies = [ - ":scala", - ], -) diff --git a/util-benchmark/src/main/scala/BUILD.bazel b/util-benchmark/src/main/scala/BUILD.bazel new file mode 100644 index 0000000000..135fde91f3 --- /dev/null +++ b/util-benchmark/src/main/scala/BUILD.bazel @@ -0,0 +1,22 @@ +scala_benchmark_jmh( + name = "jmh", + sources = [ + "**/*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/openjdk/jmh:jmh-core", + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-benchmark/src/main/java/com/twitter/json", + "util/util-core:scala", + "util/util-hashing/src/main/scala", + "util/util-jackson/src/main/scala/com/twitter/util/jackson", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-stats/src/main/scala", + "util/util-validator/src/main/scala/com/twitter/util/validation", + ], +) diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncQueueBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncQueueBenchmark.scala index f62f8ae009..2627939158 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncQueueBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncQueueBenchmark.scala @@ -9,6 +9,11 @@ class AsyncQueueBenchmark extends StdBenchAnnotations { private[this] val q = new AsyncQueue[String] + @Benchmark + def allocation(): AsyncQueue[String] = { + new AsyncQueue[String] + } + @Benchmark def offersThenPolls: Future[String] = { q.offer("hello") diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncSemaphoreBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncSemaphoreBenchmark.scala index c3cb7a0b8b..7a234eb9b5 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncSemaphoreBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncSemaphoreBenchmark.scala @@ -8,10 +8,22 @@ import org.openjdk.jmh.annotations._ @Threads(10) class AsyncSemaphoreBenchmark extends StdBenchAnnotations { - private[this] val semaphore: AsyncSemaphore = new AsyncSemaphore(1) + private[this] val noWaitersSemaphore: AsyncSemaphore = new AsyncSemaphore(10) + private[this] val waitersSemaphore: AsyncSemaphore = new AsyncSemaphore(1) + private[this] val mixedSemaphore: AsyncSemaphore = new AsyncSemaphore(5) @Benchmark - def acquireAndRelease(): Unit = { - Await.result(semaphore.acquire().map(_.release())) + def noWaiters(): Unit = { + Await.result(noWaitersSemaphore.acquire().map(_.release())) + } + + @Benchmark + def waiters(): Unit = { + Await.result(waitersSemaphore.acquire().map(_.release())) + } + + @Benchmark + def mixed(): Unit = { + Await.result(mixedSemaphore.acquire().map(_.release())) } } diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncStreamBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncStreamBenchmark.scala index d0dfc9844d..c90d5a8f3a 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncStreamBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/AsyncStreamBenchmark.scala @@ -1,7 +1,7 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ -import com.twitter.util.{Await, StdBenchAnnotations} +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Await, Future, StdBenchAnnotations} import org.openjdk.jmh.annotations._ @State(Scope.Benchmark) @@ -11,6 +11,14 @@ class AsyncStreamBenchmark extends StdBenchAnnotations { @Param(Array("10")) var size: Int = _ + /** Implementation */ + @Param(Array("cons", "embed")) + var impl: String = _ + + /** Concurrency Level */ + @Param(Array("2")) + var concurrency: Int = _ + private[this] val timeout = 5.seconds private[this] var as: AsyncStream[Int] = _ @@ -25,7 +33,14 @@ class AsyncStreamBenchmark extends StdBenchAnnotations { @Setup(Level.Iteration) def setup(): Unit = { - as = AsyncStream.fromSeq(0.until(size)) + as = impl match { + case "cons" => + AsyncStream.fromSeq(0.until(size)) + case "embed" => + 0.until(size).foldRight(AsyncStream.empty[Int]) { (i, tail) => + AsyncStream.embed(Future.value(i +:: tail)) + } + } longs = genLongStream(1000) } @@ -47,6 +62,13 @@ class AsyncStreamBenchmark extends StdBenchAnnotations { def flatMap(): Seq[Int] = Await.result(as.flatMap(FlatMapFn).toSeq()) + private[this] val MapConcurrentFn: Int => Future[Int] = + x => Future.value(x + 1) + + @Benchmark + def mapConcurrent(): Seq[Int] = + Await.result(as.mapConcurrent(concurrency)(MapConcurrentFn).toSeq()) + private[this] val FilterFn: Int => Boolean = x => x % 2 == 0 @@ -64,4 +86,4 @@ class AsyncStreamBenchmark extends StdBenchAnnotations { @Benchmark def ++(): Int = Await.result((longs ++ as).foldLeft(0)(_ + _)) -} +} \ No newline at end of file diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/BrokerBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/BrokerBenchmark.scala new file mode 100644 index 0000000000..8dd4ca4782 --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/BrokerBenchmark.scala @@ -0,0 +1,51 @@ +package com.twitter.concurrent + +import com.twitter.util.StdBenchAnnotations +import org.openjdk.jmh.annotations.{Benchmark, Scope, State} +import org.openjdk.jmh.infra.Blackhole + +@State(Scope.Benchmark) +class BrokerBenchmark extends StdBenchAnnotations { + // It's safe to share brokers across iterations because each benchmark pairs + // a send with a recv. + private[this] val q = new Broker[Int] + private[this] val p = new Broker[Int] + + @Benchmark + def sendAndRecv(blackhole: Blackhole): Unit = { + val f = q.send(1).sync() + blackhole.consume(f) + val g = q.recv.sync() + blackhole.consume(g) + } + + @Benchmark + def recvAndSend(blackhole: Blackhole): Unit = { + val g = q.recv.sync() + blackhole.consume(g) + val f = q.send(2).sync() + blackhole.consume(f) + } + + // Especially notable is that, when selecting over multiple brokers (as we do + // below), both send and recv operations are guaranteed to affect the same + // broker. Offer.select commits to the communication that can proceed, which + // means that if we attempt to send to multiple brokers, only one broker will + // be modified, the state of the other brokers is unaffected. + + @Benchmark + def selectRecvAndSend(blackhole: Blackhole): Unit = { + val f = Offer.select(q.recv, p.recv) + blackhole.consume(f) + val g = Offer.select(q.send(3), p.send(3)) + blackhole.consume(g) + } + + @Benchmark + def selectSendAndRecv(blackhole: Blackhole): Unit = { + val g = Offer.select(q.send(3), p.send(3)) + blackhole.consume(g) + val f = Offer.select(q.recv, p.recv) + blackhole.consume(f) + } +} diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/ConduitSpscBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/ConduitSpscBenchmark.scala index 5c2f175bbf..a3ebdde41e 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/ConduitSpscBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/ConduitSpscBenchmark.scala @@ -83,9 +83,7 @@ class ConduitSpscBenchmark extends StdBenchAnnotations { val stream = mkAsyncStream(size) // Consume - stream.foldLeftF(false) { (_, buf) => - sink(buf) - } + stream.foldLeftF(false) { (_, buf) => sink(buf) } }) private[this] def feedBroker(n: Int): Future[Unit] = @@ -111,8 +109,8 @@ class ConduitSpscBenchmark extends StdBenchAnnotations { else source().flatMap(r.write) before feedReader(n - 1) private[this] def consumeReader(n: Int): Future[Boolean] = - if (n <= 1) r.read(Int.MaxValue).flatMap(sink) - else r.read(Int.MaxValue).flatMap(sink).flatMap(_ => consumeReader(n - 1)) + if (n <= 1) r.read().flatMap(sink) + else r.read().flatMap(sink).flatMap(_ => consumeReader(n - 1)) @Benchmark def reader: Boolean = @@ -131,11 +129,7 @@ class ConduitSpscBenchmark extends StdBenchAnnotations { private[this] def consumeSpool[A](spool: Spool[A], b: Boolean): Future[Boolean] = if (spool.isEmpty) Future.value(b) else - sink(spool.head).flatMap { newB => - spool.tail.flatMap { tail => - consumeSpool(tail, newB) - } - } + sink(spool.head).flatMap { newB => spool.tail.flatMap { tail => consumeSpool(tail, newB) } } @Benchmark def spool: Boolean = @@ -144,8 +138,6 @@ class ConduitSpscBenchmark extends StdBenchAnnotations { val f = mkSpool(size) // Consume - f.flatMap { s => - consumeSpool(s, false) - } + f.flatMap { s => consumeSpool(s, false) } }) } diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/OfferBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/OfferBenchmark.scala index cc0c7eb6e6..1e9e952616 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/OfferBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/OfferBenchmark.scala @@ -22,7 +22,7 @@ object OfferBenchmark { // to focus on the benchmarking the impl of `choose`, use dumbed-down impls // of the offer and tx classes. private class SimpleOffer[T](fut: Future[Tx[T]]) extends Offer[T] { - override def prepare() = fut + override def prepare(): Future[Tx[T]] = fut } private class SimpleTx[T](t: T) extends Tx[T] { @@ -36,13 +36,13 @@ object OfferBenchmark { @Param(Array("3", "10", "100")) var numToChooseFrom: Int = 3 - val rng = new Random(3102957159L) - val rngOpt = Some(rng) + val rng: Random = new Random(3102957159L) + val rngOpt: Some[Random] = Some(rng) // N.B. Traditionally you'd put this type of setup inside of a benchmark setup method per invocation // but because the runtime of each invocation is so small, the overhead is high JMH advises against // using Invocation setup methods: - // http://hg.openjdk.java.net/code-tools/jmh/file/bdfc7d3a6ebf/jmh-core/src/main/java/org/openjdk/jmh/annotations/Level.java#l50 + // https://hg.openjdk.java.net/code-tools/jmh/file/bdfc7d3a6ebf/jmh-core/src/main/java/org/openjdk/jmh/annotations/Level.java#l50 def toChooseFrom(): Seq[Offer[Int]] = { val ps = new ArrayBuffer[Promise[Tx[Int]]](numToChooseFrom) val ofs = new ArrayBuffer[Offer[Int]](numToChooseFrom) @@ -56,7 +56,7 @@ object OfferBenchmark { ps(rng.nextInt(numToChooseFrom)).setValue(Tx.const(5)) - ofs + ofs.toSeq } } } diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/OnceBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/OnceBenchmark.scala index 86b0424a9f..1d1d2945bc 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/OnceBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/OnceBenchmark.scala @@ -17,6 +17,6 @@ class OnceBenchmark { object OnceBenchmark { @State(Scope.Benchmark) class OnceState { - val once = Once(()) + val once: () => Unit = Once(()) } } diff --git a/util-benchmark/src/main/scala/com/twitter/concurrent/SchedulerBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/concurrent/SchedulerBenchmark.scala index c60e380fde..775c837060 100644 --- a/util-benchmark/src/main/scala/com/twitter/concurrent/SchedulerBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/concurrent/SchedulerBenchmark.scala @@ -5,7 +5,7 @@ import org.openjdk.jmh.annotations._ import org.openjdk.jmh.infra.Blackhole // Run this via: -// ./sbt 'project util-benchmark' 'jmh:run SchedulerBenchmark -i 10 -wi 5 -f 1 -t 8 -bm avgt -tu ns' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'SchedulerBenchmark' -i 10 -wi 5 -f 1 -t 8 -bm avgt -tu ns // // Notes: // - threads, -t, should be 2x number of logical cores. diff --git a/util-benchmark/src/main/scala/com/twitter/conversions/U64Benchmark.scala b/util-benchmark/src/main/scala/com/twitter/conversions/U64Benchmark.scala new file mode 100644 index 0000000000..e3244cb9c0 --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/conversions/U64Benchmark.scala @@ -0,0 +1,33 @@ +package com.twitter.conversions + +import com.twitter.conversions.U64Ops._ +import com.twitter.util.StdBenchAnnotations +import org.openjdk.jmh.annotations._ + +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'U64Benchmark' +@State(Scope.Benchmark) +class U64Benchmark extends StdBenchAnnotations { + private[this] val rng = new scala.util.Random(42) + private[this] val hexes = Seq.fill(97)(rng.nextLong()) :+ 0L :+ Long.MaxValue :+ Long.MinValue + private[this] val longToHex: Long => String = { long: Long => long.toU64HexString } + private[this] val javaLongToHex: Long => String = (java.lang.Long.toHexString _) + private[this] val leadingZerosJavaLongToHex: Long => String = { long: Long => + "%16x".format(long) + } + + @Benchmark + def u64Hexes: Seq[String] = { + hexes.map(longToHex) + } + + @Benchmark + def javaHexes: Seq[String] = { + hexes.map(javaLongToHex) + } + + @Benchmark + def leadingZerosJavaHexes: Seq[String] = { + hexes.map(leadingZerosJavaLongToHex) + } + +} diff --git a/util-benchmark/src/main/scala/com/twitter/finagle/stats/CumulativeGaugeBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/finagle/stats/CumulativeGaugeBenchmark.scala index 04b6ca10ba..bf16f546f4 100644 --- a/util-benchmark/src/main/scala/com/twitter/finagle/stats/CumulativeGaugeBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/finagle/stats/CumulativeGaugeBenchmark.scala @@ -5,7 +5,7 @@ import java.util import java.util.concurrent.TimeUnit import org.openjdk.jmh.annotations._ -// ./sbt 'project util-benchmark' 'jmh:run CumulativeGaugeBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'CumulativeGaugeBenchmark' @State(Scope.Benchmark) class CumulativeGaugeBenchmark extends StdBenchAnnotations { import CumulativeGaugeBenchmark._ @@ -24,9 +24,7 @@ class CumulativeGaugeBenchmark extends StdBenchAnnotations { def setup(): Unit = { cg = new CGauge() gauges.clear() - 0.until(num).foreach { _ => - gauges.add(cg.addGauge(1)) - } + 0.until(num).foreach { _ => gauges.add(cg.addGauge(1, NoMetadata)) } } @Benchmark @@ -40,7 +38,7 @@ class CumulativeGaugeBenchmark extends StdBenchAnnotations { @OutputTimeUnit(TimeUnit.MICROSECONDS) @Fork(5) def addGauge(): Gauge = { - cg.addGauge(1) + cg.addGauge(1, NoMetadata) } } diff --git a/util-benchmark/src/main/scala/com/twitter/finagle/stats/StatBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/finagle/stats/StatBenchmark.scala index 405eeefb42..a48aa52478 100644 --- a/util-benchmark/src/main/scala/com/twitter/finagle/stats/StatBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/finagle/stats/StatBenchmark.scala @@ -1,10 +1,14 @@ package com.twitter.finagle.stats -import com.twitter.util.{Await, Future, StdBenchAnnotations} +import com.twitter.util.Await +import com.twitter.util.Future +import com.twitter.util.StdBenchAnnotations import java.util.concurrent.TimeUnit -import org.openjdk.jmh.annotations.{Benchmark, Scope, State} +import org.openjdk.jmh.annotations.Benchmark +import org.openjdk.jmh.annotations.Scope +import org.openjdk.jmh.annotations.State -// ./sbt 'project util-benchmark' 'jmh:run StatBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'StatBenchmark' @State(Scope.Benchmark) class StatBenchmark extends StdBenchAnnotations { diff --git a/util-benchmark/src/main/scala/com/twitter/hashing/KetamaDistributorBench.scala b/util-benchmark/src/main/scala/com/twitter/hashing/KetamaDistributorBench.scala index 53933b348b..cdb2f31bd1 100644 --- a/util-benchmark/src/main/scala/com/twitter/hashing/KetamaDistributorBench.scala +++ b/util-benchmark/src/main/scala/com/twitter/hashing/KetamaDistributorBench.scala @@ -16,13 +16,13 @@ class KetamaDistributorBench extends StdBenchAnnotations { val numReps = 640 val numHosts = 310 - val hosts = 0.until(numHosts).map { i => - "abcd-zzz-" + i + "-ab2.qwer.twitter.com" - } - val nodes = hosts map { h => - KetamaNode[String](h + ":12021", 1, h) + val hosts: scala.collection.immutable.IndexedSeq[String] = + 0.until(numHosts).map { i => "abcd-zzz-" + i + "-ab2.qwer.twitter.com" } + val nodes: scala.collection.immutable.IndexedSeq[HashNode[String]] = hosts map { h => + HashNode[String](h + ":12021", 1, h) } @Benchmark - def ketama(): KetamaDistributor[String] = new KetamaDistributor(nodes, numReps, false) + def ketama(): ConsistentHashingDistributor[String] = + new ConsistentHashingDistributor(nodes, numReps, false) } diff --git a/util-benchmark/src/main/scala/com/twitter/io/BufBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/io/BufBenchmark.scala index 3957c0c913..0ac4e3d445 100644 --- a/util-benchmark/src/main/scala/com/twitter/io/BufBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/io/BufBenchmark.scala @@ -6,7 +6,7 @@ import org.openjdk.jmh.annotations._ import org.openjdk.jmh.infra.Blackhole import scala.util.Random -// ./sbt 'project util-benchmark' 'jmh:run BufBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'BufBenchmark' @State(Scope.Benchmark) class BufBenchmark extends StdBenchAnnotations { diff --git a/util-benchmark/src/main/scala/com/twitter/io/BufConcatBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/io/BufConcatBenchmark.scala index 6fe0d215ed..659005f53a 100644 --- a/util-benchmark/src/main/scala/com/twitter/io/BufConcatBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/io/BufConcatBenchmark.scala @@ -30,10 +30,22 @@ object BufConcatBenchmark { val buf4: Buf = buf3.concat(buf1) val bufMatrix: Array[(Buf, Buf)] = Array( - (buf1, buf1), (buf1, buf2), (buf1, buf3), (buf1, buf4), - (buf2, buf1), (buf2, buf2), (buf2, buf3), (buf2, buf4), - (buf3, buf1), (buf3, buf2), (buf3, buf3), (buf3, buf4), - (buf4, buf1), (buf4, buf2), (buf4, buf3), (buf4, buf4) + (buf1, buf1), + (buf1, buf2), + (buf1, buf3), + (buf1, buf4), + (buf2, buf1), + (buf2, buf2), + (buf2, buf3), + (buf2, buf4), + (buf3, buf1), + (buf3, buf2), + (buf3, buf3), + (buf3, buf4), + (buf4, buf1), + (buf4, buf2), + (buf4, buf3), + (buf4, buf4) ) @State(Scope.Thread) @@ -69,11 +81,9 @@ class BufConcatBenchmark extends StdBenchAnnotations { class ConcatTwoBufsBenchmark extends StdBenchAnnotations { import BufConcatBenchmark._ - @Param(Array( - "0", "1", "2", "3", - "4", "5", "6", "7", - "8", "9", "10", "11", - "12", "13", "14", "15")) + @Param( + Array("0", "1", "2", "3", "4", "5", "6", "7", "8", "9", "10", "11", "12", "13", "14", "15") + ) var index: Int = _ @Benchmark diff --git a/util-benchmark/src/main/scala/com/twitter/io/PipeBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/io/PipeBenchmark.scala new file mode 100644 index 0000000000..27d0e7d68c --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/io/PipeBenchmark.scala @@ -0,0 +1,32 @@ +package com.twitter.io + +import com.twitter.util.Future +import com.twitter.util.StdBenchAnnotations +import org.openjdk.jmh.annotations.Benchmark +import org.openjdk.jmh.annotations.Measurement +import org.openjdk.jmh.annotations.Scope +import org.openjdk.jmh.annotations.State +import org.openjdk.jmh.annotations.Warmup + +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'PipeBenchmark' +@State(Scope.Benchmark) +@Warmup(iterations = 2) +@Measurement(iterations = 2) +class PipeBenchmark extends StdBenchAnnotations { + + private val foo = Buf.Utf8("foo") + + @Benchmark + def writeAndRead(): Future[Option[Buf]] = { + val p = new Pipe[Buf] + p.write(foo) + p.read() + } + + @Benchmark + def readAndWrite(): Future[Unit] = { + val p = new Pipe[Buf] + p.read() + p.write(foo) + } +} diff --git a/util-benchmark/src/main/scala/com/twitter/json/JacksonBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/json/JacksonBenchmark.scala new file mode 100644 index 0000000000..a2d956b934 --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/json/JacksonBenchmark.scala @@ -0,0 +1,75 @@ +package com.twitter.json + +import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} +import com.twitter.util.StdBenchAnnotations +import com.twitter.util.jackson.ScalaObjectMapper +import java.io.ByteArrayInputStream +import java.time.LocalDateTime +import org.openjdk.jmh.annotations._ + +object JacksonBenchmark { + case class TestTask(request_id: String, group_ids: Seq[String], params: TestTaskParams) + + case class TestTaskParams( + results: TaskTaskResults, + start_time: LocalDateTime, + end_time: LocalDateTime, + priority: String) + + case class TaskTaskResults( + country_codes: Seq[String], + group_by: Seq[String], + format: Option[TestFormat], + demographics: Seq[TestDemographic]) + +} + +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'JacksonBenchmark' +@State(Scope.Benchmark) +class JacksonBenchmark extends StdBenchAnnotations { + import JacksonBenchmark._ + + private[this] val json = + """{ + "request_id": "00000000-1111-2222-3333-444444444444", + "type": "impression", + "params": { + "priority": "normal", + "start_time": "2013-01-01T00:02:00.000Z", + "end_time": "2013-01-01T00:03:00.000Z", + "results": { + "demographics": ["age", "gender"], + "country_codes": ["US"], + "group_by": ["group"] + } + }, + "group_ids":["grp-1"] + }""" + + private[this] val bytes = json.getBytes("UTF-8") + + /** create framework scala object mapper with all defaults */ + private[this] val mapperWithCaseClassDeserializer: ScalaObjectMapper = + ScalaObjectMapper.builder.objectMapper + + /** create framework scala object mapper without case class deserializer */ + private[this] val mapperWithoutCaseClassDeserializer: ScalaObjectMapper = { + // configure the underlying Jackson ScalaObjectMapper with everything but the CaseClassJacksonModule + val underlying = new com.fasterxml.jackson.databind.ObjectMapper with JacksonScalaObjectMapper + underlying.registerModules(ScalaObjectMapper.DefaultJacksonModules: _*) + // do not use apply which will always install the default jackson modules on the underlying mapper + new ScalaObjectMapper(underlying) + } + + @Benchmark + def withCaseClassDeserializer(): TestTask = { + val is = new ByteArrayInputStream(bytes) + mapperWithCaseClassDeserializer.parse[TestTask](is) + } + + @Benchmark + def withoutCaseClassDeserializer(): TestTask = { + val is = new ByteArrayInputStream(bytes) + mapperWithoutCaseClassDeserializer.parse[TestTask](is) + } +} diff --git a/util-benchmark/src/main/scala/com/twitter/jvm/CpuProfile.scala b/util-benchmark/src/main/scala/com/twitter/jvm/CpuProfile.scala index 710a4beb57..808da933c7 100644 --- a/util-benchmark/src/main/scala/com/twitter/jvm/CpuProfile.scala +++ b/util-benchmark/src/main/scala/com/twitter/jvm/CpuProfile.scala @@ -4,6 +4,7 @@ import java.lang.management.ManagementFactory import java.util.concurrent.TimeUnit import org.openjdk.jmh.annotations._ import scala.util.Random +import java.lang.management.ThreadMXBean @OutputTimeUnit(TimeUnit.NANOSECONDS) @BenchmarkMode(Array(Mode.AverageTime)) @@ -30,12 +31,12 @@ class CpuProfileBenchmark { object CpuProfileBenchmark { @State(Scope.Benchmark) class ThreadState { - val bean = ManagementFactory.getThreadMXBean() + val bean: ThreadMXBean = ManagementFactory.getThreadMXBean() // TODO: change dynamically. bounce up & down // the stack. μ and σ are for *this* stack. class Stack(rng: Random, μ: Int, σ: Int) { - def apply() = { + def apply(): Int = { val depth = math.max(1, μ + (rng.nextGaussian * σ).toInt) dive(depth) } @@ -59,9 +60,9 @@ object CpuProfileBenchmark { @Setup(Level.Iteration) def setUp(): Unit = { val stack = new Stack(new Random(rngSeed), stackMeanSize, stackStddev) - threads = for (_ <- 0 until nthreads) - yield - new Thread { + threads = + for (_ <- 0 until nthreads) + yield new Thread { override def run(): Unit = { try stack() catch { @@ -70,19 +71,13 @@ object CpuProfileBenchmark { } } - threads foreach { t => - t.start() - } + threads foreach { t => t.start() } } @TearDown(Level.Iteration) def tearDown(): Unit = { - threads foreach { t => - t.interrupt() - } - threads foreach { t => - t.join() - } + threads foreach { t => t.interrupt() } + threads foreach { t => t.join() } threads = Seq() } var threads = Seq[Thread]() diff --git a/util-benchmark/src/main/scala/com/twitter/string/StringConcatenationBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/string/StringConcatenationBenchmark.scala index 8e908dc244..e97359859b 100644 --- a/util-benchmark/src/main/scala/com/twitter/string/StringConcatenationBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/string/StringConcatenationBenchmark.scala @@ -4,14 +4,14 @@ import com.twitter.util.StdBenchAnnotations import org.openjdk.jmh.annotations._ import scala.util.Random -// ./sbt 'project util-benchmark' 'jmh:run StringConcatenationBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'StringConcatenationBenchmark' @State(Scope.Benchmark) class StringConcatenationBenchmark extends StdBenchAnnotations { val Length = 16 val N = 10000 - val rng = new Random(1010101) - val words = rng.alphanumeric.grouped(Length).map(_.mkString).take(N).toArray + val rng: Random = new Random(1010101) + val words: Array[String] = rng.alphanumeric.grouped(Length).map(_.mkString).take(N).toArray var i = 0 private[this] def word(): String = { diff --git a/util-benchmark/src/main/scala/com/twitter/util/DurationBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/DurationBenchmark.scala index 70e461de18..c27bd5886d 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/DurationBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/DurationBenchmark.scala @@ -2,7 +2,7 @@ package com.twitter.util import org.openjdk.jmh.annotations._ -// ./sbt 'project util-benchmark' 'jmh:run DurationBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'DurationBenchmark' @State(Scope.Benchmark) class DurationBenchmark extends StdBenchAnnotations { diff --git a/util-benchmark/src/main/scala/com/twitter/util/FutureBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/FutureBenchmark.scala index 02967a1f9a..8c6ffcf945 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/FutureBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/FutureBenchmark.scala @@ -3,6 +3,7 @@ package com.twitter.util import com.twitter.concurrent.Offer import org.openjdk.jmh.annotations._ +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'FutureBenchmark' class FutureBenchmark extends StdBenchAnnotations { import FutureBenchmark._ @@ -11,6 +12,14 @@ class FutureBenchmark extends StdBenchAnnotations { def timePromise(): Future[Unit] = new Promise[Unit] + @Benchmark + def timeConstUnit(): Future[Unit] = + Future.Done.unit + + @Benchmark + def timeConstVoid(): Future[Void] = + Future.Done.voided + @Benchmark def timeBy(state: ByState): Future[Unit] = { import state._ @@ -76,6 +85,22 @@ class FutureBenchmark extends StdBenchAnnotations { new Promise[Unit].flatMap(FlatMapFn) } + @Benchmark + def fromConstMap: Future[Unit] = + Future.Done.map(MapFn) + + @Benchmark + def fromConstFlatMapToConst: Future[Unit] = + Future.Done.flatMap(FlatMapFn) + + @Benchmark + def fromConstTransformToConst: Future[Unit] = + Future.Done.transform(TransformFn) + + @Benchmark + def timeConstFilter(): Future[Unit] = + Future.Done.filter(FilterFn) + @Benchmark def timeSelect(state: SelectState): Future[(Try[Unit], Seq[Future[Unit]])] = { import state._ @@ -198,6 +223,37 @@ class FutureBenchmark extends StdBenchAnnotations { Await.result(p) } + def buildChain(root: Promise[Int], state: RunqState): Future[Int] = { + var next: Future[Int] = root + var sum = 0 + var i = 0 + + val fmap = { i: Int => Future.value(i + 1) } + val respond = { _: Try[Int] => sum += 1 } + val map = { i: Int => i + 1 } + val filter = { _: Int => true } + + while (i < state.depth) { + next = i % 4 match { + case 0 => next.flatMap(fmap) + case 1 => next.respond(respond) + case 2 => next.map(map) + case 3 => next.filter(filter) + } + i += 1 + } + next + } + + @Benchmark + def runChain(state: RunqState): Int = { + val p = new Promise[Int]() + val f = buildChain(p, state) + // trigger the callbacks + p.setValue(5) + Await.result(f) + } + @Benchmark def monitored(): String = { val f = Future.monitored { @@ -216,21 +272,17 @@ class FutureBenchmark extends StdBenchAnnotations { object FutureBenchmark { final val N = 10 - private val RespondFn: Try[Unit] => Unit = { _ => - () - } - private val FlatMapFn: Unit => Future[Unit] = { _: Unit => - Future.Unit - } - private val MapFn: Unit => Unit = { _: Unit => - () - } + private val RespondFn: Try[Unit] => Unit = { _ => () } + private val FlatMapFn: Unit => Future[Unit] = { _: Unit => Future.Unit } + private val MapFn: Unit => Unit = { _: Unit => () } + private val TransformFn: Try[Unit] => Future[Unit] = { _: Try[Unit] => Future.Done } + private val FilterFn: Unit => Boolean = { _ => true } @State(Scope.Benchmark) class ByState { - val timer = Timer.Nil + val timer: Timer = Timer.Nil val now: Time = Time.now - val exc = new TimeoutException("") + val exc: TimeoutException = new TimeoutException("") } @State(Scope.Benchmark) @@ -250,9 +302,7 @@ object FutureBenchmark { @Setup def prepare(): Unit = { - futures = (0 until size).map { i => - Future.value(i) - }.toList + futures = (0 until size).map { i => Future.value(i) }.toList } } @@ -266,12 +316,12 @@ object FutureBenchmark { @State(Scope.Benchmark) class FlattenState { - val future = Future.value(Future.Unit) + val future: Future[Future[Unit]] = Future.value(Future.Unit) } @State(Scope.Benchmark) class LowerFromTryState { - val future = Future.Unit.liftToTry + val future: Future[Try[Unit]] = Future.Unit.liftToTry } @State(Scope.Benchmark) @@ -286,7 +336,7 @@ object FutureBenchmark { @Param(Array("0", "1", "10", "100")) var numToSelect = 0 - val p = new Promise[Unit] + val p: Promise[Unit] = new Promise[Unit] val futures: Seq[Future[Unit]] = if (numToSelect == 0) IndexedSeq.empty @@ -300,7 +350,7 @@ object FutureBenchmark { @Param(Array("0", "1", "10", "100")) var numToSelect = 0 - val p = new Promise[Unit] + val p: Promise[Unit] = new Promise[Unit] val futures: IndexedSeq[Future[Unit]] = if (numToSelect == 0) IndexedSeq.empty diff --git a/util-benchmark/src/main/scala/com/twitter/util/LocalBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/LocalBenchmark.scala index 7b66ab8416..2950cb09fa 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/LocalBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/LocalBenchmark.scala @@ -2,12 +2,14 @@ package com.twitter.util import org.openjdk.jmh.annotations._ -@State(Scope.Benchmark) +@State(Scope.Thread) +@Threads(1) class LocalBenchmark extends StdBenchAnnotations { // We've sampled some service instances and get the result that usually // 10 Locals are constructed and 3 Locals are set. val size = 10 private[this] val locals = Array.fill(size)(new Local[String]) + private[this] val getter = locals(1).threadLocalGetter() @Benchmark def let1(): String = @@ -41,4 +43,7 @@ class LocalBenchmark extends StdBenchAnnotations { @Benchmark def get(): Option[String] = locals(1)() + + @Benchmark + def getThreadLocal(): Option[String] = getter() } diff --git a/util-benchmark/src/main/scala/com/twitter/util/MonitorBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/MonitorBenchmark.scala index f04b6c8ae6..7fd71a1061 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/MonitorBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/MonitorBenchmark.scala @@ -1,9 +1,11 @@ package com.twitter.util -import org.openjdk.jmh.annotations.{Benchmark, Scope, State} +import org.openjdk.jmh.annotations.Benchmark +import org.openjdk.jmh.annotations.Scope +import org.openjdk.jmh.annotations.State import scala.util.control.NoStackTrace -// ./sbt 'project util-benchmark' 'jmh:run MonitorBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'MonitorBenchmark' @State(Scope.Benchmark) class MonitorBenchmark extends StdBenchAnnotations { diff --git a/util-benchmark/src/main/scala/com/twitter/util/PromiseBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/PromiseBenchmark.scala index ebce037ca8..4291556d7e 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/PromiseBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/PromiseBenchmark.scala @@ -4,11 +4,26 @@ import com.twitter.util.Promise.Detachable import org.openjdk.jmh.annotations._ @State(Scope.Benchmark) +@Warmup(iterations = 5) +@Measurement(iterations = 5) class PromiseBenchmark extends StdBenchAnnotations { + import PromiseBenchmark._ - private[this] val StringFuture = Future.value("hi") + @Benchmark + def responder(state: ResponderState): Int = { + import state._ + initPromise.updateIfEmpty(UpdateValue) + Await.result(promise) + } - private[this] val Value = Return("okok") + @Benchmark + def responderWithResourceUsage(state: ResourceTrackedResponderState): Int = { + import state._ + initPromise.updateIfEmpty(UpdateValue) + Await.result(promise) + } + + private[this] val NoopCallback: Try[String] => Unit = _ => () @Benchmark def attached(): Promise[String] = { @@ -16,11 +31,17 @@ class PromiseBenchmark extends StdBenchAnnotations { } @Benchmark - def detach(state: PromiseBenchmark.PromiseDetachState): Unit = { + def detach(state: PromiseDetachState): Unit = { import state._ promise.detach() } + @Benchmark + def detachDetachableFuture(state: PromiseDetachState): Unit = { + import state._ + detachableFuture.detach() + } + // used to isolate the work in the `updateIfEmpty` benchmark @Benchmark def newUnsatisfiedPromise(): Promise[String] = { @@ -34,22 +55,22 @@ class PromiseBenchmark extends StdBenchAnnotations { } @Benchmark - def interrupts1(state: PromiseBenchmark.InterruptsState): Promise[String] = { + def interrupts1(state: InterruptsState): Promise[String] = { Promise.interrupts(state.a) } @Benchmark - def interrupts2(state: PromiseBenchmark.InterruptsState): Promise[String] = { + def interrupts2(state: InterruptsState): Promise[String] = { Promise.interrupts(state.a, state.b) } @Benchmark - def interruptsN(state: PromiseBenchmark.InterruptsState): Promise[String] = { + def interruptsN(state: InterruptsState): Promise[String] = { Promise.interrupts(state.futures: _*) } @Benchmark - def promiseIsDefined(state: PromiseBenchmark.PromiseState): Boolean = { + def promiseIsDefined(state: PromiseState): Boolean = { if (state.i == state.len) state.i = 0 val ret = state.promises(state.i).isDefined state.i += 1 @@ -57,15 +78,64 @@ class PromiseBenchmark extends StdBenchAnnotations { } @Benchmark - def promisePoll(state: PromiseBenchmark.PromiseState): Option[Try[Unit]] = { + def promisePoll(state: PromiseState): Option[Try[Unit]] = { if (state.i == state.len) state.i = 0 val ret = state.promises(state.i).poll state.i += 1 ret } + + @Benchmark + def become(state: PromiseBenchmark.PromiseState): Promise[String] = { + val p = new Promise[String](state.f) + p.respond(NoopCallback) + + val q = new Promise[String] + + q.become(p) // p.link(q) + p + } } object PromiseBenchmark { + final val StringFuture = Future.value("hi") + final val Value = Return("okok") + + @State(Scope.Thread) + class ResponderState { + val UpdateValue = Return(1) + + @Param(Array("1", "10", "100", "1000")) + var numK: Int = _ + + var initPromise: Promise[Int] = _ + var promise: Future[Int] = _ + + @Setup(Level.Invocation) + def prepare(): Unit = { + initPromise = new Promise[Int]() + promise = (0 to numK).foldLeft[Future[Int]](initPromise) { (p, i) => + p.map(num => num + i) + } + } + } + + @State(Scope.Thread) + class ResourceTrackedResponderState extends ResponderState { + val resourceTracker = new ResourceTracker() + + @Setup(Level.Invocation) + override def prepare(): Unit = { + ResourceTracker.set(resourceTracker) + super.prepare() + } + + @TearDown(Level.Invocation) + def tearDown(): Unit = { + ResourceTracker.clear() + } + } + @State(Scope.Thread) class PromiseDetachState { @@ -76,11 +146,17 @@ object PromiseBenchmark { var attachedFutures: List[Future[Unit]] = _ var promise: Promise[Unit] with Detachable = _ + var detachableFuture: Promise[Int] with Detachable = _ + @Setup def prepare(): Unit = { global = new Promise[Unit] attachedFutures = List.fill(numAttached) { Promise.attached(global) } promise = Promise.attached(global) + + // explicitly using a Future that is not a `Promise` to + // test the `DetachableFuture` code paths. + detachableFuture = Promise.attached(new ConstFuture(Return(5))) } } diff --git a/util-benchmark/src/main/scala/com/twitter/util/StopwatchBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/StopwatchBenchmark.scala index 5f2ae94381..3a600bf53b 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/StopwatchBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/StopwatchBenchmark.scala @@ -17,11 +17,21 @@ class StopwatchBenchmark { def timeTime(state: StopwatchState): Duration = { state.elapsed() } + + @Benchmark + def timeSystemNanos(): Long = { + Stopwatch.systemNanos() + } + + @Benchmark + def timeTimeNanos(): Long = { + Stopwatch.timeNanos() + } } object StopwatchBenchmark { @State(Scope.Benchmark) class StopwatchState { - val elapsed = Stopwatch.start() + val elapsed: Stopwatch.Elapsed = Stopwatch.start() } } diff --git a/util-benchmark/src/main/scala/com/twitter/util/TimeBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/TimeBenchmark.scala index c71cc06611..d65d89b4de 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/TimeBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/TimeBenchmark.scala @@ -2,7 +2,7 @@ package com.twitter.util import org.openjdk.jmh.annotations._ -// ./sbt 'project util-benchmark' 'jmh:run TimeBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'TimeBenchmark' @State(Scope.Benchmark) class TimeBenchmark extends StdBenchAnnotations { @@ -32,4 +32,10 @@ class TimeBenchmark extends StdBenchAnnotations { @Benchmark def timeDiffOverflow: Duration = t5.diff(t4) + + @Benchmark + def timeNow: Time = Time.now + + @Benchmark + def nowNanoPrecision: Time = Time.nowNanoPrecision } diff --git a/util-benchmark/src/main/scala/com/twitter/util/TimerBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/TimerBenchmark.scala index 2f3f22b4c6..f2490a520d 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/TimerBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/TimerBenchmark.scala @@ -1,6 +1,6 @@ package com.twitter.util -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import java.util.concurrent.TimeUnit import org.openjdk.jmh.annotations._ @@ -9,7 +9,7 @@ import org.openjdk.jmh.annotations._ * performance of future scheduling (particularly because this * does not wait for the tasks to be executed). */ -// ./sbt 'project util-benchmark' 'jmh:run TimerBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'TimerBenchmark' @State(Scope.Benchmark) @Warmup(batchSize = 250) @Measurement(batchSize = 1000) diff --git a/util-benchmark/src/main/scala/com/twitter/util/TryBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/TryBenchmark.scala index 339c444fb1..62472f16a0 100644 --- a/util-benchmark/src/main/scala/com/twitter/util/TryBenchmark.scala +++ b/util-benchmark/src/main/scala/com/twitter/util/TryBenchmark.scala @@ -2,7 +2,7 @@ package com.twitter.util import org.openjdk.jmh.annotations._ -// ./sbt 'project util-benchmark' 'jmh:run TryBenchmark' +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'TryBenchmark' @State(Scope.Benchmark) class TryBenchmark extends StdBenchAnnotations { diff --git a/util-benchmark/src/main/scala/com/twitter/util/TypesBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/TypesBenchmark.scala new file mode 100644 index 0000000000..69170913fc --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/util/TypesBenchmark.scala @@ -0,0 +1,56 @@ +package com.twitter.util + +import com.twitter.util.reflect.Types +import org.openjdk.jmh.annotations._ +import scala.util.control.NonFatal + +private object TypesBenchmark { + class NotCaseClass { + val a: String = "A" + val b: String = "B" + } + + case class IsCaseClass(a: String, b: String) +} + +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'TypesBenchmark' +@State(Scope.Benchmark) +class TypesBenchmark extends StdBenchAnnotations { + import TypesBenchmark._ + + @Benchmark + def notCaseClassWithClass(): Unit = { + try { + Types.notCaseClass(classOf[NotCaseClass]) + } catch { + case NonFatal(_) => // avoid throwing exceptions so the benchmark can finish + } + } + + @Benchmark + def notCaseClassWithCaseClass(): Unit = { + try { + Types.notCaseClass(classOf[IsCaseClass]) + } catch { + case NonFatal(_) => // avoid throwing exceptions so the benchmark can finish + } + } + + @Benchmark + def caseClassWithClass(): Unit = { + try { + Types.isCaseClass(classOf[NotCaseClass]) + } catch { + case NonFatal(_) => // avoid throwing exceptions so the benchmark can finish + } + } + + @Benchmark + def caseClassWithCaseClass(): Unit = { + try { + Types.isCaseClass(classOf[IsCaseClass]) + } catch { + case NonFatal(_) => // avoid throwing exceptions so the benchmark can finish + } + } +} diff --git a/util-benchmark/src/main/scala/com/twitter/util/ValidationBenchmark.scala b/util-benchmark/src/main/scala/com/twitter/util/ValidationBenchmark.scala new file mode 100644 index 0000000000..70c51004ad --- /dev/null +++ b/util-benchmark/src/main/scala/com/twitter/util/ValidationBenchmark.scala @@ -0,0 +1,83 @@ +package com.twitter.util + +import com.twitter.util.validation.MethodValidation +import com.twitter.util.validation.ScalaValidator +import com.twitter.util.validation.engine.MethodValidationResult +import jakarta.validation.ValidationException +import jakarta.validation.constraints.Min +import jakarta.validation.constraints.NotEmpty +import org.openjdk.jmh.annotations.Benchmark +import org.openjdk.jmh.annotations.Scope +import org.openjdk.jmh.annotations.State +import scala.util.control.NonFatal + +private object ValidationBenchmark { + case class User(@NotEmpty id: String, name: String, @Min(18) age: Int) + + case class Users(@NotEmpty users: Seq[User]) { + @MethodValidation + def uniqueUsers: MethodValidationResult = + MethodValidationResult.validIfTrue( + users.map(_.id).distinct.size == users.size, + "user ids are not distinct.") + } +} + +// ./bazel run //util/util-benchmark/src/main/scala:jmh -- 'ValidationBenchmark' +@State(Scope.Benchmark) +class ValidationBenchmark extends StdBenchAnnotations { + import ValidationBenchmark._ + + private[this] val validator: ScalaValidator = ScalaValidator() + private[this] val validUser: User = User("1234567", "jack", 21) + private[this] val invalidUser: User = User("", "notJack", 13) + private[this] val nestedValidUser: Users = Users(Seq(validUser)) + private[this] val nestedInvalidUser: Users = Users(Seq(invalidUser)) + private[this] val nestedDuplicateUser: Users = Users(Seq(validUser, validUser)) + + @Benchmark + def withValidUser(): Unit = { + validator.verify(validUser) + } + + @Benchmark + def withInvalidUser(): Unit = { + try { + validator.verify(invalidUser) + } catch { + case _: ValidationException => // avoid throwing exceptions so the benchmark can finish + } + } + + @Benchmark + def withNestedValidUser(): Unit = { + validator.verify(nestedValidUser) + } + + @Benchmark + def withNestedInvalidUser(): Unit = { + try { + validator.verify(nestedInvalidUser) + } catch { + case _: ValidationException => // avoid throwing exceptions so the benchmark can finish + } + } + + @Benchmark + def withNestedDuplicateUser(): Unit = { + try { + validator.verify(nestedDuplicateUser) + } catch { + case _: ValidationException => // avoid throwing exceptions so the benchmark can finish + } + } + + override def finalize(): Unit = { + try { + validator.close() + } catch { + case NonFatal(_) => // do nothing + } + super.finalize() + } +} diff --git a/util-cache-guava/src/main/scala/BUILD b/util-cache-guava/src/main/scala/BUILD index 84fdbe5475..07a15b9193 100644 --- a/util-cache-guava/src/main/scala/BUILD +++ b/util-cache-guava/src/main/scala/BUILD @@ -1,14 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-cache-guava", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/google/guava", - "util/util-cache/src/main/scala", - "util/util-core/src/main/scala", + "util/util-cache-guava/src/main/scala/com/twitter/cache/guava", ], ) diff --git a/util-cache-guava/src/main/scala/com/twitter/cache/guava/BUILD b/util-cache-guava/src/main/scala/com/twitter/cache/guava/BUILD new file mode 100644 index 0000000000..2acd27445a --- /dev/null +++ b/util-cache-guava/src/main/scala/com/twitter/cache/guava/BUILD @@ -0,0 +1,16 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-cache-guava", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/google/guava", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-core:util-core-util", + ], +) diff --git a/util-cache-guava/src/main/scala/com/twitter/cache/guava/GuavaCache.scala b/util-cache-guava/src/main/scala/com/twitter/cache/guava/GuavaCache.scala index 1f98ebbd5a..4a0f468c1f 100644 --- a/util-cache-guava/src/main/scala/com/twitter/cache/guava/GuavaCache.scala +++ b/util-cache-guava/src/main/scala/com/twitter/cache/guava/GuavaCache.scala @@ -6,30 +6,32 @@ import com.twitter.util.Future import java.util.concurrent.Callable /** - * A [[com.twitter.cache.FutureCache]] backed by a - * [[com.google.common.cache.Cache]]. + * A `com.twitter.cache.FutureCache` backed by a + * `com.google.common.cache.Cache`. * * Any correct implementation should make sure that you evict failed results, * and don't interrupt the underlying request that has been fired off. - * [[EvictingCache]] and [[Future#interrupting]] are useful tools for building + * [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are useful tools for building * correct FutureCaches. A reference implementation for caching the results of * an asynchronous function with a guava Cache can be found at * [[GuavaCache$.fromCache]]. */ class GuavaCache[K, V](cache: GCache[K, Future[V]]) extends ConcurrentMapCache[K, V](cache.asMap) { override def getOrElseUpdate(k: K)(v: => Future[V]): Future[V] = - cache.get(k, new Callable[Future[V]] { - def call(): Future[V] = v - }) + cache.get( + k, + new Callable[Future[V]] { + def call(): Future[V] = v + }) } /** * A [[com.twitter.cache.FutureCache]] backed by a - * [[com.google.common.cache.LoadingCache]]. + * `com.google.common.cache.LoadingCache`. * Any correct implementation should make sure that you evict failed results, * and don't interrupt the underlying request that has been fired off. - * [[EvictingCache]] and [[Future#interrupting]] are useful tools for building + * [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are useful tools for building * correct FutureCaches. A reference implementation for caching the results of * an asynchronous function with a guava LoadingCache can be found at * [[GuavaCache$.fromLoadingCache]]. @@ -47,18 +49,16 @@ object GuavaCache { /** * Creates a function which properly handles the asynchronous behavior of - * [[com.google.common.cache.LoadingCache]]. + * `com.google.common.cache.LoadingCache`. */ def fromLoadingCache[K, V](cache: LoadingCache[K, Future[V]]): K => Future[V] = { val evicting = EvictingCache.lazily(new LoadingFutureCache(cache)); - { key: K => - evicting.get(key).get.interruptible() - } + { (key: K) => evicting.get(key).get.interruptible() } } /** * Creates a function which caches the results of `fn` in a - * [[com.google.common.cache.Cache]]. + * `com.google.common.cache.Cache`. */ def fromCache[K, V](fn: K => Future[V], cache: GCache[K, Future[V]]): K => Future[V] = FutureCache.default(fn, new GuavaCache(cache)) diff --git a/util-cache-guava/src/test/java/BUILD b/util-cache-guava/src/test/java/BUILD deleted file mode 100644 index f41f341c77..0000000000 --- a/util-cache-guava/src/test/java/BUILD +++ /dev/null @@ -1,13 +0,0 @@ -junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/com/google/guava", - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-cache", - "util/util-cache-guava/src/main/scala", - "util/util-core/src/main/scala", - ], -) diff --git a/util-cache-guava/src/test/java/BUILD.bazel b/util-cache-guava/src/test/java/BUILD.bazel new file mode 100644 index 0000000000..c922f72e7b --- /dev/null +++ b/util-cache-guava/src/test/java/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-cache-guava/src/test/java/com/twitter/cache/guava", + ], +) diff --git a/util-cache-guava/src/test/java/com/twitter/cache/guava/BUILD.bazel b/util-cache-guava/src/test/java/com/twitter/cache/guava/BUILD.bazel new file mode 100644 index 0000000000..aea988d311 --- /dev/null +++ b/util-cache-guava/src/test/java/com/twitter/cache/guava/BUILD.bazel @@ -0,0 +1,19 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/com/google/guava", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-cache-guava/src/main/scala/com/twitter/cache/guava", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-core:util-core-util", + ], +) diff --git a/util-cache-guava/src/test/scala/BUILD b/util-cache-guava/src/test/scala/BUILD deleted file mode 100644 index 7dc39a0789..0000000000 --- a/util-cache-guava/src/test/scala/BUILD +++ /dev/null @@ -1,15 +0,0 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/com/google/code/findbugs:jsr305", - "3rdparty/jvm/com/google/guava", - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-cache-guava/src/main/scala", - "util/util-cache/src/main/scala", - "util/util-cache/src/test/scala:abstract_tests", - "util/util-core/src/main/scala", - ], -) diff --git a/util-cache-guava/src/test/scala/BUILD.bazel b/util-cache-guava/src/test/scala/BUILD.bazel new file mode 100644 index 0000000000..4970d3153b --- /dev/null +++ b/util-cache-guava/src/test/scala/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-cache-guava/src/test/scala/com/twitter/cache/guava", + ], +) diff --git a/util-cache-guava/src/test/scala/com/twitter/cache/guava/BUILD.bazel b/util-cache-guava/src/test/scala/com/twitter/cache/guava/BUILD.bazel new file mode 100644 index 0000000000..cdc325f6e2 --- /dev/null +++ b/util-cache-guava/src/test/scala/com/twitter/cache/guava/BUILD.bazel @@ -0,0 +1,21 @@ +junit_tests( + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/com/google/code/findbugs:jsr305", + "3rdparty/jvm/com/google/guava", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-cache-guava/src/main/scala/com/twitter/cache/guava", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-cache/src/test/scala/com/twitter/cache:abstract_tests", + "util/util-core:util-core-util", + ], +) diff --git a/util-cache-guava/src/test/scala/com/twitter/cache/guava/GuavaCacheTest.scala b/util-cache-guava/src/test/scala/com/twitter/cache/guava/GuavaCacheTest.scala index 675fa956c0..13612f638a 100644 --- a/util-cache-guava/src/test/scala/com/twitter/cache/guava/GuavaCacheTest.scala +++ b/util-cache-guava/src/test/scala/com/twitter/cache/guava/GuavaCacheTest.scala @@ -3,10 +3,7 @@ package com.twitter.cache.guava import com.google.common.cache.{CacheLoader, CacheBuilder} import com.twitter.cache.AbstractFutureCacheTest import com.twitter.util.{Future, Promise} -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class GuavaCacheTest extends AbstractFutureCacheTest { def name: String = "GuavaCache" diff --git a/util-cache-guava/src/test/scala/com/twitter/cache/guava/LoadingFutureCacheTest.scala b/util-cache-guava/src/test/scala/com/twitter/cache/guava/LoadingFutureCacheTest.scala index 3be8ea6cac..4e1c7a5371 100644 --- a/util-cache-guava/src/test/scala/com/twitter/cache/guava/LoadingFutureCacheTest.scala +++ b/util-cache-guava/src/test/scala/com/twitter/cache/guava/LoadingFutureCacheTest.scala @@ -3,10 +3,7 @@ package com.twitter.cache.guava import com.google.common.cache.{CacheLoader, CacheBuilder} import com.twitter.cache.AbstractLoadingFutureCacheTest import com.twitter.util.Future -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class LoadingFutureCacheTest extends AbstractLoadingFutureCacheTest { def name: String = "LoadingFutureCache (guava)" @@ -15,7 +12,7 @@ class LoadingFutureCacheTest extends AbstractLoadingFutureCacheTest { // loading cache semantics are sufficiently unique // to merit distinct tests. - def mkCtx: Ctx = new Ctx { + def mkCtx(): Ctx = new Ctx { val cache = new LoadingFutureCache( CacheBuilder .newBuilder() diff --git a/util-cache/BUILD b/util-cache/BUILD index ce5424c315..2efa01e029 100644 --- a/util-cache/BUILD +++ b/util-cache/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-cache/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-cache/src/test/java", "util/util-cache/src/test/scala", diff --git a/util-cache/src/main/scala/BUILD b/util-cache/src/main/scala/BUILD index f1e5b11ca8..adfc85a118 100644 --- a/util-cache/src/main/scala/BUILD +++ b/util-cache/src/main/scala/BUILD @@ -1,16 +1,22 @@ scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, + # This target exists only to act as an aggregator jar. + sources = ["PantsWorkaroundCache.scala"], + platform = "java8", provides = scala_artifact( org = "com.twitter", name = "util-cache", repo = artifactory, ), + strict_deps = True, + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", - "util/util-core/src/main/scala", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-cache/src/main/scala/com/twitter/cache/caffeine", + "util/util-core:util-core-util", ], exports = [ - "util/util-core/src/main/scala", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-cache/src/main/scala/com/twitter/cache/caffeine", + "util/util-core:util-core-util", ], ) diff --git a/util-cache/src/main/scala/PantsWorkaroundCache.scala b/util-cache/src/main/scala/PantsWorkaroundCache.scala new file mode 100644 index 0000000000..2dccfb94cd --- /dev/null +++ b/util-cache/src/main/scala/PantsWorkaroundCache.scala @@ -0,0 +1,4 @@ +// workaround for https://github.com/pantsbuild/pants/issues/6678 +private final class PantsWorkaroundCache { + throw new IllegalStateException("workaround for https://github.com/pantsbuild/pants/issues/6678") +} diff --git a/util-cache/src/main/scala/com/twitter/cache/BUILD b/util-cache/src/main/scala/com/twitter/cache/BUILD new file mode 100644 index 0000000000..4e526fc10a --- /dev/null +++ b/util-cache/src/main/scala/com/twitter/cache/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-cache-cache", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-cache/src/main/scala/com/twitter/cache/ConcurrentMapCache.scala b/util-cache/src/main/scala/com/twitter/cache/ConcurrentMapCache.scala index ba14ccb4ad..e1f3c7806c 100644 --- a/util-cache/src/main/scala/com/twitter/cache/ConcurrentMapCache.scala +++ b/util-cache/src/main/scala/com/twitter/cache/ConcurrentMapCache.scala @@ -5,14 +5,14 @@ import java.util.concurrent.ConcurrentMap /** * A [[com.twitter.cache.FutureCache]] backed by a - * [[java.util.concurrent.ConcurrentMap]] + * `java.util.concurrent.ConcurrentMap` * * Any correct implementation should make sure that you evict failed * results, and don't interrupt the underlying request that has been - * fired off. [[EvictingCache]] and [[Future#interrupting]] are + * fired off. [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are * useful tools for building correct FutureCaches. A reference * implementation for caching the results of an asynchronous function - * with a [[ConcurrentMap]] can be found at [[FutureCache$.fromMap]]. + * with a `ConcurrentMap` can be found at [[FutureCache$.fromMap]]. */ class ConcurrentMapCache[K, V](underlying: ConcurrentMap[K, Future[V]]) extends FutureCache[K, V] { def get(key: K): Option[Future[V]] = Option(underlying.get(key)) @@ -20,7 +20,7 @@ class ConcurrentMapCache[K, V](underlying: ConcurrentMap[K, Future[V]]) extends def set(key: K, value: Future[V]): Unit = underlying.put(key, value) def getOrElseUpdate(key: K)(compute: => Future[V]): Future[V] = { - val p = Promise[V] + val p = Promise[V]() underlying.putIfAbsent(key, p) match { case null => p.become(compute) diff --git a/util-cache/src/main/scala/com/twitter/cache/EvictingCache.scala b/util-cache/src/main/scala/com/twitter/cache/EvictingCache.scala index d56ee28cb8..220b6c0907 100644 --- a/util-cache/src/main/scala/com/twitter/cache/EvictingCache.scala +++ b/util-cache/src/main/scala/com/twitter/cache/EvictingCache.scala @@ -5,9 +5,7 @@ import com.twitter.util.{Future, Throw} private[cache] class EvictingCache[K, V](underlying: FutureCache[K, V]) extends FutureCacheProxy[K, V](underlying) { private[this] def evictOnFailure(k: K, f: Future[V]): Future[V] = { - f.onFailure { _ => - evict(k, f) - } + f.onFailure { _ => evict(k, f) } f // we return the original future to make evict(k, f) easier to work with. } @@ -17,9 +15,11 @@ private[cache] class EvictingCache[K, V](underlying: FutureCache[K, V]) } override def getOrElseUpdate(k: K)(v: => Future[V]): Future[V] = - evictOnFailure(k, underlying.getOrElseUpdate(k) { - v - }) + evictOnFailure( + k, + underlying.getOrElseUpdate(k) { + v + }) } private[cache] class LazilyEvictingCache[K, V](underlying: FutureCache[K, V]) diff --git a/util-cache/src/main/scala/com/twitter/cache/FutureCache.scala b/util-cache/src/main/scala/com/twitter/cache/FutureCache.scala index 43a8bc1571..83e74116d2 100644 --- a/util-cache/src/main/scala/com/twitter/cache/FutureCache.scala +++ b/util-cache/src/main/scala/com/twitter/cache/FutureCache.scala @@ -10,7 +10,7 @@ import java.util.concurrent.ConcurrentMap * * Any correct implementation should make sure that you evict failed * results, and don't interrupt the underlying request that has been - * fired off. [[EvictingCache]] and [[Future#interrupting]] are + * fired off. [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are * useful tools for building correct FutureCaches. A reference * implementation for caching the results of an asynchronous function * can be found at [[FutureCache$.default]]. @@ -87,8 +87,8 @@ abstract class FutureCacheProxy[K, V](underlying: FutureCache[K, V]) extends Fut object FutureCache { /** - * A [[com.twitter.cache.FutureCache]] backed by a - * [[java.util.concurrent.ConcurrentMap]]. + * A `com.twitter.cache.FutureCache` backed by a + * `java.util.concurrent.ConcurrentMap`. */ def fromMap[K, V](fn: K => Future[V], map: ConcurrentMap[K, Future[V]]): K => Future[V] = default(fn, new ConcurrentMapCache(map)) @@ -107,9 +107,7 @@ object FutureCache { * @see [[standard]] for the equivalent Java API. */ def default[K, V](fn: K => Future[V], cache: FutureCache[K, V]): K => Future[V] = - AsyncMemoize(fn, new EvictingCache(cache)) andThen { f: Future[V] => - f.interruptible() - } + AsyncMemoize(fn, new EvictingCache(cache)) andThen { (f: Future[V]) => f.interruptible() } /** * Alias for [[default]] which can be called from Java. diff --git a/util-cache/src/main/scala/com/twitter/cache/KeyEncodingCache.scala b/util-cache/src/main/scala/com/twitter/cache/KeyEncodingCache.scala index 05a2c693e9..e25ee05c14 100644 --- a/util-cache/src/main/scala/com/twitter/cache/KeyEncodingCache.scala +++ b/util-cache/src/main/scala/com/twitter/cache/KeyEncodingCache.scala @@ -20,10 +20,8 @@ import com.twitter.util.Future * { case (target: Int, _, _, _) => target } * )) */ -private[cache] class KeyEncodingCache[K, V, U]( - encode: K => V, - underlying: FutureCache[V, U] -) extends FutureCache[K, U] { +private[cache] class KeyEncodingCache[K, V, U](encode: K => V, underlying: FutureCache[V, U]) + extends FutureCache[K, U] { override def get(key: K): Option[Future[U]] = underlying.get(encode(key)) def set(key: K, value: Future[U]): Unit = underlying.set(encode(key), value) diff --git a/util-cache/src/main/scala/com/twitter/cache/Refresh.scala b/util-cache/src/main/scala/com/twitter/cache/Refresh.scala index 261cd03a6b..9c13a0c7ee 100644 --- a/util-cache/src/main/scala/com/twitter/cache/Refresh.scala +++ b/util-cache/src/main/scala/com/twitter/cache/Refresh.scala @@ -17,7 +17,7 @@ import java.util.concurrent.atomic.AtomicReference * def getData(): Future[T] = { ... } * * can be memoized with a TTL of 1 hour as follows: - * import com.twitter.util.TimeConversions._ + * import com.twitter.conversions.DurationOps._ * import com.twitter.cache.Refresh * val getData: () => Future[T] = Refresh.every(1.hour) { ... } */ @@ -33,7 +33,7 @@ object Refresh { * * From Java: * Refresh.every(Duration.fromSeconds(3600), new Function0>() { - * @Override + * \@Override * public Future apply() { * ... * } @@ -65,6 +65,6 @@ object Refresh { } } // Return result, which is a no-arg function that returns a future - result + result _ } } diff --git a/util-cache/src/main/scala/com/twitter/cache/caffeine/BUILD b/util-cache/src/main/scala/com/twitter/cache/caffeine/BUILD new file mode 100644 index 0000000000..cb1bc43396 --- /dev/null +++ b/util-cache/src/main/scala/com/twitter/cache/caffeine/BUILD @@ -0,0 +1,16 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-cache-caffeine", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-core:util-core-util", + ], +) diff --git a/util-cache/src/main/scala/com/twitter/cache/caffeine/CaffeineCache.scala b/util-cache/src/main/scala/com/twitter/cache/caffeine/CaffeineCache.scala index 75c7226e39..0a853d4949 100644 --- a/util-cache/src/main/scala/com/twitter/cache/caffeine/CaffeineCache.scala +++ b/util-cache/src/main/scala/com/twitter/cache/caffeine/CaffeineCache.scala @@ -11,7 +11,7 @@ import java.util.function.Function * * Any correct implementation should make sure that you evict failed results, * and don't interrupt the underlying request that has been fired off. - * [[EvictingCache]] and [[Future#interrupting]] are useful tools for building + * [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are useful tools for building * correct FutureCaches. A reference implementation for caching the results of * an asynchronous function with a caffeine Cache can be found at * [[CaffeineCache$.fromCache]]. @@ -19,9 +19,11 @@ import java.util.function.Function class CaffeineCache[K, V](cache: Cache[K, Future[V]]) extends ConcurrentMapCache[K, V](cache.asMap) { override def getOrElseUpdate(k: K)(v: => Future[V]): Future[V] = - cache.get(k, new Function[K, Future[V]] { - def apply(k: K): Future[V] = v - }) + cache.get( + k, + new Function[K, Future[V]] { + def apply(k: K): Future[V] = v + }) } /** @@ -30,7 +32,7 @@ class CaffeineCache[K, V](cache: Cache[K, Future[V]]) * Any correct implementation should make sure that you evict failed results, * and don't interrupt the underlying request that has been fired off. - * [[EvictingCache]] and [[Future#interrupting]] are useful tools for building + * [[EvictingCache$]] and interrupting [[com.twitter.util.Future]]s are useful tools for building * correct FutureCaches. A reference implementation for caching the results of * an asynchronous function with a caffeine LoadingCache can be found at * [[CaffeineCache$.fromLoadingCache]]. @@ -52,9 +54,7 @@ object CaffeineCache { */ def fromLoadingCache[K, V](cache: LoadingCache[K, Future[V]]): K => Future[V] = { val evicting = EvictingCache.lazily(new LoadingFutureCache(cache)); - { key: K => - evicting.get(key).get.interruptible() - } + { (key: K) => evicting.get(key).get.interruptible() } } /** diff --git a/util-cache/src/test/java/BUILD b/util-cache/src/test/java/BUILD index cd15d7536c..a5cf81040f 100644 --- a/util-cache/src/test/java/BUILD +++ b/util-cache/src/test/java/BUILD @@ -1,12 +1,18 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ - "3rdparty/jvm/com/google/guava", "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/mockito:mockito-core", "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-cache", - "util/util-core/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-core:util-core-util", ], ) diff --git a/util-cache/src/test/scala/BUILD b/util-cache/src/test/scala/BUILD index e2adcffa8a..99b826c8d9 100644 --- a/util-cache/src/test/scala/BUILD +++ b/util-cache/src/test/scala/BUILD @@ -1,28 +1,7 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-cache", - "util/util-core/src/main/scala", - ], -) - -scala_library( - name = "abstract_tests", - sources = [ - "com/twitter/cache/AbstractFutureCacheTest.scala", - "com/twitter/cache/AbstractLoadingFutureCacheTest.scala", - ], - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-cache", - "util/util-core/src/main/scala", + "util/util-cache/src/test/scala/com/twitter/cache", + "util/util-cache/src/test/scala/com/twitter/cache/caffeine", ], ) diff --git a/util-cache/src/test/scala/com/twitter/cache/AbstractFutureCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/AbstractFutureCacheTest.scala index d3b3e9fed2..07b3036984 100644 --- a/util-cache/src/test/scala/com/twitter/cache/AbstractFutureCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/AbstractFutureCacheTest.scala @@ -1,10 +1,10 @@ package com.twitter.cache import com.twitter.util.{Future, Promise} -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite // TODO: should also check for races -abstract class AbstractFutureCacheTest extends FunSuite { +abstract class AbstractFutureCacheTest extends AnyFunSuite { def name: String diff --git a/util-cache/src/test/scala/com/twitter/cache/AbstractLoadingFutureCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/AbstractLoadingFutureCacheTest.scala index 03a4334ff7..27fd175b1a 100644 --- a/util-cache/src/test/scala/com/twitter/cache/AbstractLoadingFutureCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/AbstractLoadingFutureCacheTest.scala @@ -1,10 +1,10 @@ package com.twitter.cache -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, Future} -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -abstract class AbstractLoadingFutureCacheTest extends FunSuite { +abstract class AbstractLoadingFutureCacheTest extends AnyFunSuite { // NB we can't reuse AbstractFutureCacheTest since // loading cache semantics are sufficiently unique // to merit distinct tests. diff --git a/util-cache/src/test/scala/com/twitter/cache/BUILD b/util-cache/src/test/scala/com/twitter/cache/BUILD new file mode 100644 index 0000000000..1f24e52bac --- /dev/null +++ b/util-cache/src/test/scala/com/twitter/cache/BUILD @@ -0,0 +1,46 @@ +junit_tests( + sources = [ + "!AbstractFutureCacheTest.scala", + "!AbstractLoadingFutureCacheTest.scala", + "*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + ":abstract_tests", + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-cache/src/main/scala/com/twitter/cache/caffeine", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + ], +) + +scala_library( + name = "abstract_tests", + sources = [ + "AbstractFutureCacheTest.scala", + "AbstractLoadingFutureCacheTest.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + ], +) diff --git a/util-cache/src/test/scala/com/twitter/cache/ConcurrentMapCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/ConcurrentMapCacheTest.scala index 5d3c3ea6b1..1462b3dc17 100644 --- a/util-cache/src/test/scala/com/twitter/cache/ConcurrentMapCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/ConcurrentMapCacheTest.scala @@ -2,10 +2,7 @@ package com.twitter.cache import com.twitter.util.Future import java.util.concurrent.ConcurrentHashMap -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class ConcurrentMapCacheTest extends AbstractFutureCacheTest { def name: String = "ConcurrentMapCache" diff --git a/util-cache/src/test/scala/com/twitter/cache/EvictingCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/EvictingCacheTest.scala index 23b97e55fd..c7cf5f3ccb 100644 --- a/util-cache/src/test/scala/com/twitter/cache/EvictingCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/EvictingCacheTest.scala @@ -2,18 +2,15 @@ package com.twitter.cache import com.twitter.util.{Promise, Future} import java.util.concurrent.ConcurrentHashMap -import org.junit.runner.RunWith import org.mockito.Mockito.{verify, never} -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class EvictingCacheTest extends FunSuite with MockitoSugar { +class EvictingCacheTest extends AnyFunSuite with MockitoSugar { test("EvictingCache should evict on failed futures for set") { val cache = mock[FutureCache[String, String]] val fCache = new EvictingCache(cache) - val p = Promise[String] + val p = Promise[String]() fCache.set("key", p) verify(cache).set("key", p) p.setException(new Exception) @@ -23,7 +20,7 @@ class EvictingCacheTest extends FunSuite with MockitoSugar { test("EvictingCache should keep satisfied futures for set") { val cache = mock[FutureCache[String, String]] val fCache = new EvictingCache(cache) - val p = Promise[String] + val p = Promise[String]() fCache.set("key", p) verify(cache).set("key", p) p.setValue("value") @@ -34,7 +31,7 @@ class EvictingCacheTest extends FunSuite with MockitoSugar { val map = new ConcurrentHashMap[String, Future[String]]() val cache = new ConcurrentMapCache(map) val fCache = new EvictingCache(cache) - val p = Promise[String] + val p = Promise[String]() assert(fCache.getOrElseUpdate("key")(p).poll == p.poll) p.setException(new Exception) assert(fCache.get("key") == None) @@ -44,7 +41,7 @@ class EvictingCacheTest extends FunSuite with MockitoSugar { val map = new ConcurrentHashMap[String, Future[String]]() val cache = new ConcurrentMapCache(map) val fCache = new EvictingCache(cache) - val p = Promise[String] + val p = Promise[String]() assert(fCache.getOrElseUpdate("key")(p).poll == p.poll) p.setValue("value") assert(fCache.get("key").map(_.poll) == Some(p.poll)) diff --git a/util-cache/src/test/scala/com/twitter/cache/KeyEncodingCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/KeyEncodingCacheTest.scala index 5c696d40de..3f7a296aef 100644 --- a/util-cache/src/test/scala/com/twitter/cache/KeyEncodingCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/KeyEncodingCacheTest.scala @@ -2,10 +2,7 @@ package com.twitter.cache import java.util.concurrent.ConcurrentHashMap import com.twitter.util.Future -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class KeyEncodingCacheTest extends AbstractFutureCacheTest { def name: String = "KeyEncodingCache" @@ -13,8 +10,6 @@ class KeyEncodingCacheTest extends AbstractFutureCacheTest { val underlyingMap: ConcurrentHashMap[Int, Future[String]] = new ConcurrentHashMap() val underlyingCache: FutureCache[Int, String] = new ConcurrentMapCache(underlyingMap) val cache: FutureCache[String, String] = - new KeyEncodingCache({ num: String => - num.hashCode - }, underlyingCache) + new KeyEncodingCache({ (num: String) => num.hashCode }, underlyingCache) } } diff --git a/util-cache/src/test/scala/com/twitter/cache/LazilyEvictingCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/LazilyEvictingCacheTest.scala index 5e74c224eb..8cbce28c25 100644 --- a/util-cache/src/test/scala/com/twitter/cache/LazilyEvictingCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/LazilyEvictingCacheTest.scala @@ -3,12 +3,12 @@ package com.twitter.cache import com.github.benmanes.caffeine.cache.{CacheLoader, Caffeine} import com.twitter.cache.caffeine.LoadingFutureCache import com.twitter.util.{Await, Future, Promise} -import org.mockito.Matchers._ +import org.mockito.ArgumentMatchers._ import org.mockito.Mockito.{never, verify, when} -import org.scalatest.FunSuite -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -class LazilyEvictingCacheTest extends FunSuite with MockitoSugar { +class LazilyEvictingCacheTest extends AnyFunSuite with MockitoSugar { private val explodingCacheLoader = new CacheLoader[String, Future[String]] { def load(k: String): Future[String] = @@ -85,14 +85,16 @@ class LazilyEvictingCacheTest extends FunSuite with MockitoSugar { var loadCount = 0 val cache = new LoadingFutureCache( - Caffeine.newBuilder().build( - new CacheLoader[String, Future[Int]] { - def load(k: String): Future[Int] = { - loadCount += 1 - Future.value(loadCount) + Caffeine + .newBuilder() + .build( + new CacheLoader[String, Future[Int]] { + def load(k: String): Future[Int] = { + loadCount += 1 + Future.value(loadCount) + } } - } - ) + ) ) val fCache = new LazilyEvictingCache(cache) @@ -117,14 +119,16 @@ class LazilyEvictingCacheTest extends FunSuite with MockitoSugar { var loadCount = 0 val cache = new LoadingFutureCache( - Caffeine.newBuilder().build( - new CacheLoader[String, Future[Int]] { - def load(k: String): Future[Int] = { - loadCount += 1 - Future.value(loadCount) + Caffeine + .newBuilder() + .build( + new CacheLoader[String, Future[Int]] { + def load(k: String): Future[Int] = { + loadCount += 1 + Future.value(loadCount) + } } - } - ) + ) ) val fCache = new LazilyEvictingCache(cache) diff --git a/util-cache/src/test/scala/com/twitter/cache/RefreshTest.scala b/util-cache/src/test/scala/com/twitter/cache/RefreshTest.scala index 4bb33121d6..bf872583dd 100644 --- a/util-cache/src/test/scala/com/twitter/cache/RefreshTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/RefreshTest.scala @@ -1,15 +1,12 @@ package com.twitter.cache -import com.twitter.util.{Await, Future, Time, Promise} +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Await, Future, Promise, Time} import org.mockito.Mockito._ -import com.twitter.util.TimeConversions._ -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest._ -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class RefreshTest extends FunSuite with MockitoSugar { +class RefreshTest extends AnyFunSuite with MockitoSugar { class Ctx { val provider = mock[() => Future[Int]] @@ -73,7 +70,7 @@ class RefreshTest extends FunSuite with MockitoSugar { val ctx = new Ctx import ctx._ - val promise = Promise[Int] + val promise = Promise[Int]() reset(provider) when(provider()) .thenReturn(promise) @@ -93,7 +90,7 @@ class RefreshTest extends FunSuite with MockitoSugar { val ctx = new Ctx import ctx._ - val promise = Promise[Int] + val promise = Promise[Int]() reset(provider) when(provider()).thenReturn(promise) @@ -108,7 +105,7 @@ class RefreshTest extends FunSuite with MockitoSugar { verify(provider, times(1))() reset(provider) - val promise2 = Promise[Int] + val promise2 = Promise[Int]() when(provider()).thenReturn(promise2) val result3 = memoizedFuture() diff --git a/util-cache/src/test/scala/com/twitter/cache/caffeine/BUILD b/util-cache/src/test/scala/com/twitter/cache/caffeine/BUILD new file mode 100644 index 0000000000..36a1d33fa3 --- /dev/null +++ b/util-cache/src/test/scala/com/twitter/cache/caffeine/BUILD @@ -0,0 +1,21 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/com/google/code/findbugs:jsr305", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-cache/src/main/scala/com/twitter/cache", + "util/util-cache/src/main/scala/com/twitter/cache/caffeine", + "util/util-cache/src/test/scala/com/twitter/cache:abstract_tests", + "util/util-core:util-core-util", + ], +) diff --git a/util-cache/src/test/scala/com/twitter/cache/caffeine/CaffeineCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/caffeine/CaffeineCacheTest.scala index 0a9dfeaff0..19ec88a095 100644 --- a/util-cache/src/test/scala/com/twitter/cache/caffeine/CaffeineCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/caffeine/CaffeineCacheTest.scala @@ -3,10 +3,7 @@ package com.twitter.cache.caffeine import com.github.benmanes.caffeine.cache.{Caffeine, CacheLoader, LoadingCache} import com.twitter.cache.AbstractFutureCacheTest import com.twitter.util.{Future, Promise} -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class CaffeineCacheTest extends AbstractFutureCacheTest { def name: String = "CaffeineCache" diff --git a/util-cache/src/test/scala/com/twitter/cache/caffeine/LoadingFutureCacheTest.scala b/util-cache/src/test/scala/com/twitter/cache/caffeine/LoadingFutureCacheTest.scala index a13cd950f4..efe2223c15 100644 --- a/util-cache/src/test/scala/com/twitter/cache/caffeine/LoadingFutureCacheTest.scala +++ b/util-cache/src/test/scala/com/twitter/cache/caffeine/LoadingFutureCacheTest.scala @@ -3,10 +3,7 @@ package com.twitter.cache.caffeine import com.github.benmanes.caffeine.cache.{Caffeine, CacheLoader} import com.twitter.cache.AbstractLoadingFutureCacheTest import com.twitter.util.Future -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) class LoadingFutureCacheTest extends AbstractLoadingFutureCacheTest { def name: String = "LoadingFutureCache (caffeine)" @@ -15,7 +12,7 @@ class LoadingFutureCacheTest extends AbstractLoadingFutureCacheTest { // loading cache semantics are sufficiently unique // to merit distinct tests. - def mkCtx: Ctx = new Ctx { + def mkCtx(): Ctx = new Ctx { val cache = new LoadingFutureCache( Caffeine .newBuilder() diff --git a/util-codec/BUILD b/util-codec/BUILD index 9a12f6e0e4..4a62ee8b7e 100644 --- a/util-codec/BUILD +++ b/util-codec/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-codec/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-codec/src/test/scala", ], diff --git a/util-codec/OWNERS b/util-codec/OWNERS deleted file mode 100644 index c313f6c1d5..0000000000 --- a/util-codec/OWNERS +++ /dev/null @@ -1,5 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta diff --git a/util-codec/PROJECT b/util-codec/PROJECT new file mode 100644 index 0000000000..4edda67ec7 --- /dev/null +++ b/util-codec/PROJECT @@ -0,0 +1,6 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan + - mdickinson diff --git a/util-codec/src/main/scala/BUILD b/util-codec/src/main/scala/BUILD index d576eacc6b..a1251e2edf 100644 --- a/util-codec/src/main/scala/BUILD +++ b/util-codec/src/main/scala/BUILD @@ -1,12 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-codec", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "util/util-core/src/main/scala", + "util/util-codec/src/main/scala/com/twitter/util", ], ) diff --git a/util-codec/src/main/scala/com/twitter/util/BUILD b/util-codec/src/main/scala/com/twitter/util/BUILD new file mode 100644 index 0000000000..d88e57df30 --- /dev/null +++ b/util-codec/src/main/scala/com/twitter/util/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-codec", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core/src/main/scala/com/twitter/io", + ], +) diff --git a/util-codec/src/main/scala/com/twitter/util/StringEncoder.scala b/util-codec/src/main/scala/com/twitter/util/StringEncoder.scala index a7afb865b6..785abf192c 100644 --- a/util-codec/src/main/scala/com/twitter/util/StringEncoder.scala +++ b/util-codec/src/main/scala/com/twitter/util/StringEncoder.scala @@ -68,7 +68,7 @@ trait GZIPStringEncoder extends StringEncoder { Base64StringEncoder.encode(baos.toByteArray) } - def encodeString(str: String) = encode(str.getBytes("UTF-8")) + def encodeString(str: String): String = encode(str.getBytes("UTF-8")) override def decode(str: String): Array[Byte] = { val baos = new ByteArrayOutputStream diff --git a/util-codec/src/test/scala/BUILD b/util-codec/src/test/scala/BUILD index 0a758329c1..a3e7a3addd 100644 --- a/util-codec/src/test/scala/BUILD +++ b/util-codec/src/test/scala/BUILD @@ -1,10 +1,6 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-codec/src/main/scala", + "util/util-codec/src/test/scala/com/twitter/util", ], ) diff --git a/util-codec/src/test/scala/com/twitter/util/BUILD b/util-codec/src/test/scala/com/twitter/util/BUILD new file mode 100644 index 0000000000..0c1e041423 --- /dev/null +++ b/util-codec/src/test/scala/com/twitter/util/BUILD @@ -0,0 +1,16 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-codec/src/main/scala/com/twitter/util", + ], +) diff --git a/util-codec/src/test/scala/com/twitter/util/StringEncoderTest.scala b/util-codec/src/test/scala/com/twitter/util/StringEncoderTest.scala index 89f8c8b415..df60791a82 100644 --- a/util-codec/src/test/scala/com/twitter/util/StringEncoderTest.scala +++ b/util-codec/src/test/scala/com/twitter/util/StringEncoderTest.scala @@ -7,7 +7,7 @@ package com.twitter.util * 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 + * https://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, @@ -16,68 +16,59 @@ package com.twitter.util * limitations under the License. */ -import java.lang.StringBuilder +import org.scalatest.funsuite.AnyFunSuite -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class StringEncoderTest extends WordSpec { +class StringEncoderTest extends AnyFunSuite { val longString = "A string that is really really really really really really long and has more than 76 characters" val result = "QSBzdHJpbmcgdGhhdCBpcyByZWFsbHkgcmVhbGx5IHJlYWxseSByZWFsbHkgcmVhbGx5IHJlYWxseSBsb25nIGFuZCBoYXMgbW9yZSB0aGFuIDc2IGNoYXJhY3RlcnM=" - "strip new lines" in { + test("strip new lines") { assert(Base64StringEncoder.encode(longString.getBytes) == result) } - "decode value with stripped new lines" in { + test("decode value with stripped new lines") { assert(new String(Base64StringEncoder.decode(result)) == longString) } } -@RunWith(classOf[JUnitRunner]) -class Base64StringEncoderTest extends WordSpec { +class Base64StringEncoderTest extends AnyFunSuite { val urlUnsafeBytes = Array(-1.toByte, -32.toByte) val resultUnsafe = "/+A" val resultSafe = "_-A" - "encode / as _ and encode + as - to maintain url safe strings" in { + test("encode / as _ and encode + as - to maintain url safe strings") { assert(Base64UrlSafeStringEncoder.encode(urlUnsafeBytes) == resultSafe) } - "decode url-safe strings" in { + test("decode url-safe strings") { assert(Base64UrlSafeStringEncoder.decode(resultSafe) === urlUnsafeBytes) } - "decode url-unsafe strings" in { + test("decode url-unsafe strings") { intercept[IllegalArgumentException] { Base64UrlSafeStringEncoder.decode(resultUnsafe) } } } -@RunWith(classOf[JUnitRunner]) -class GZIPStringEncoderTest extends WordSpec { - "a gzip string encoder" should { +class GZIPStringEncoderTest extends AnyFunSuite { + test("a gzip string encoder should properly encode and decode strings") { val gse = new GZIPStringEncoder {} - "properly encode and decode strings" in { - def testCodec(str: String): Unit = { - assert(str == gse.decodeString(gse.encodeString(str))) - } + def testCodec(str: String): Unit = { + assert(str == gse.decodeString(gse.encodeString(str))) + } - testCodec("a") - testCodec("\n\t\n\t\n\n\n\n\t\n\nt\n\t\n\t\n\tn\t\nt\nt\nt\nt\nt\nt\tn\nt\nt\n\t\nt\n") - testCodec("aosnetuhsaontehusaonethsoantehusaonethusonethusnaotehu") + testCodec("a") + testCodec("\n\t\n\t\n\n\n\n\t\n\nt\n\t\n\t\n\tn\t\nt\nt\nt\nt\nt\nt\tn\nt\nt\n\t\nt\n") + testCodec("aosnetuhsaontehusaonethsoantehusaonethusonethusnaotehu") - // build a huge string - val sb = new StringBuilder - for (_ <- 1 to 10000) { - sb.append("oasnuthoesntihosnteidosentidosentauhsnoetidosentihsoneitdsnuthsin\n") - } - testCodec(sb.toString) + // build a huge string + val sb = new StringBuilder + for (_ <- 1 to 10000) { + sb.append("oasnuthoesntihosnteidosentidosentauhsnoetidosentihsoneitdsnuthsin\n") } + testCodec(sb.toString) } } diff --git a/util-collection/BUILD b/util-collection/BUILD deleted file mode 100644 index ef294a8e31..0000000000 --- a/util-collection/BUILD +++ /dev/null @@ -1,12 +0,0 @@ -target( - dependencies = [ - "util/util-collection/src/main/scala", - ], -) - -target( - name = "tests", - dependencies = [ - "util/util-collection/src/test/scala", - ], -) diff --git a/util-collection/src/main/scala/BUILD b/util-collection/src/main/scala/BUILD deleted file mode 100644 index 70b90ebb47..0000000000 --- a/util-collection/src/main/scala/BUILD +++ /dev/null @@ -1,15 +0,0 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-collection", - repo = artifactory, - ), - dependencies = [ - "util/util-core/src/main/scala", - ], - exports = [ - "util/util-core/src/main/scala", - ], -) diff --git a/util-collection/src/main/scala/com/twitter/collection/GenerationalQueue.scala b/util-collection/src/main/scala/com/twitter/collection/GenerationalQueue.scala deleted file mode 100644 index a14d65dac9..0000000000 --- a/util-collection/src/main/scala/com/twitter/collection/GenerationalQueue.scala +++ /dev/null @@ -1,160 +0,0 @@ -package com.twitter.collection - -import scala.collection._ - -import com.twitter.util.{Duration, Time} - -trait GenerationalQueue[A] { - def touch(a: A): Unit - def add(a: A): Unit - def remove(a: A): Unit - def collect(d: Duration): Option[A] - def collectAll(d: Duration): Iterable[A] -} - -/** - * Generational Queue keep track of elements based on their last activity. - * You can refresh activity of an element by calling touch(a: A) on it. - * There is 2 ways of retrieving old elements: - * - collect(age: Duration) collect the oldest element (age of elem must be > age) - * - collectAll(age: Duration) collect all the elements which age > age in parameter - */ -class ExactGenerationalQueue[A] extends GenerationalQueue[A] { - private[this] val container = mutable.HashMap.empty[A, Time] - private[this] implicit val ordering = Ordering.by[(A, Time), Time] { case (_, ts) => ts } - - /** - * touch insert the element if it is not yet present - */ - def touch(a: A) = synchronized { container.update(a, Time.now) } - - def add(a: A) = synchronized { container += ((a, Time.now)) } - - def remove(a: A) = synchronized { container.remove(a) } - - def collect(age: Duration): Option[A] = synchronized { - if (container.isEmpty) - None - else - container.min match { - case (a, t) if (t.untilNow > age) => Some(a) - case _ => None - } - } - - def collectAll(age: Duration): Iterable[A] = synchronized { - (container filter { case (_, t) => t.untilNow > age }).keys - } -} - -/** - * Improved GenerationalQueue: using a list of buckets responsible for containing elements belonging - * to a slice of time. - * For instance: 3 Buckets, First contains elements from 0 to 10, second elements from 11 to 20 - * and third elements from 21 to 30 - * We expand the list when we need a new bucket, and compact the list to stash all old buckets - * into one. - * There is a slightly difference with the other implementation, when we collect elements we only - * choose randomly an element in the oldest bucket, as we don't have activity date in this bucket - * we consider the worst case and then we can miss some expired elements by never find elements - * that aren't expired. - */ -class BucketGenerationalQueue[A](timeout: Duration) extends GenerationalQueue[A] { - object TimeBucket { - def empty[B] = new TimeBucket[B](Time.now, timeSlice) - } - class TimeBucket[B](var origin: Time, var span: Duration) extends mutable.HashSet[B] { - - // return the age of the potential youngest element of the bucket (may be negative if the - // bucket is not yet expired) - def age(now: Time = Time.now): Duration = (origin + span).until(now) - - def ++=(other: TimeBucket[B]) = { - origin = List(origin, other.origin).min - span = List(origin + span, other.origin + other.span).max - origin - - super.++=(other) - this - } - - override def toString() = "TimeBucket(origin=%d, size=%d, age=%s, count=%d)".format( - origin.inMilliseconds, - span.inMilliseconds, - age().toString, - super.size - ) - } - - private[this] val timeSlice = timeout / 3 - private[this] var buckets = List[TimeBucket[A]]() - - private[this] def maybeGrowChain() = { - // NB: age of youngest element is negative when bucket isn't expired - val growChain = buckets.headOption - .map((bucket) => { - bucket.age() > Duration.Zero - }) - .getOrElse(true) - - if (growChain) - buckets = TimeBucket.empty[A] :: buckets - growChain - } - - private[this] def compactChain(): List[TimeBucket[A]] = { - val now = Time.now - // partition equivalent to takeWhile/dropWhile because buckets are ordered - val (news, olds) = buckets.partition(_.age(now) < timeout) - if (olds.isEmpty) - news - else { - val tailBucket = olds.head - olds drop 1 foreach { tailBucket ++= _ } - if (tailBucket.isEmpty) - news - else - news ::: List(tailBucket) - } - } - - private[this] def updateBuckets(): Unit = { - if (maybeGrowChain()) - buckets = compactChain() - } - - def touch(a: A) = synchronized { - buckets drop 1 foreach { _.remove(a) } - add(a) - } - - def add(a: A) = synchronized { - updateBuckets() - buckets.head.add(a) - } - - def remove(a: A) = synchronized { - buckets foreach { _.remove(a) } - buckets = compactChain() - } - - def collect(d: Duration): Option[A] = synchronized { - if (buckets.isEmpty) - return None - - if (buckets.last.isEmpty) - buckets = compactChain() - - if (buckets.isEmpty) - return None - - val oldestBucket = buckets.last - if (d < oldestBucket.age()) - oldestBucket.headOption - else - None - } - - def collectAll(d: Duration): Iterable[A] = synchronized { - (buckets dropWhile (_.age() < d)).flatten - } -} diff --git a/util-collection/src/main/scala/com/twitter/collection/RecordSchema.scala b/util-collection/src/main/scala/com/twitter/collection/RecordSchema.scala deleted file mode 100644 index 6f958b2d8d..0000000000 --- a/util-collection/src/main/scala/com/twitter/collection/RecordSchema.scala +++ /dev/null @@ -1,222 +0,0 @@ -package com.twitter.collection - -import java.util.IdentityHashMap - -/** - * RecordSchema represents the declaration of a heterogeneous - * [[com.twitter.collection.RecordSchema.Record Record]] type, with - * [[com.twitter.collection.RecordSchema.Field Fields]] that are determined at runtime. A Field - * declares the static type of its associated values, so although the record itself is dynamic, - * field access is type-safe. - * - * Given a RecordSchema declaration `schema`, any number of - * [[com.twitter.collection.RecordSchema.Record Records]] of that schema can be obtained with - * `schema.newRecord`. The type that Scala assigns to this value is what Scala calls a - * "[[http://lampwww.epfl.ch/~amin/dot/fpdt.pdf path-dependent type]]," meaning that - * `schema1.Record` and `schema2.Record` name distinct types. The same is true of fields: - * `schema1.Field[A]` and `schema2.Field[A]` are distinct, and can only be used with the - * corresponding Record. - */ -final class RecordSchema { - - private[RecordSchema] final class Entry(var value: Any) { - var locked: Boolean = false - } - - /** - * Record is an instance of a [[com.twitter.collection.RecordSchema RecordSchema]] declaration. - * Records are mutable; the `update` method assigns or reassigns a value to a given field. If - * the user requires that a field's assigned value is never reassigned later, the user can `lock` - * that field. - * - * '''Note that this implementation is not synchronized.''' If multiple threads access a record - * concurrently, and at least one of the threads modifies the record, it ''must'' be synchronized - * externally. - */ - final class Record private[RecordSchema] ( - fields: IdentityHashMap[Field[_], Entry] = new IdentityHashMap[Field[_], Entry] - ) { - - private[this] def getOrInitializeEntry(field: Field[_]): Entry = { - var entry = fields.get(field) - if (entry eq null) { - entry = new Entry(field.default()) - fields.put(field, entry) - } - entry - } - - /** - * Returns the current value assigned to `field` in this record. If there is no value currently - * assigned (either explicitly by a previous `update`, or by the field's declared default), this - * will throw an `IllegalStateException`, indicating the field was never initialized. - * - * Note that Scala provides two syntactic equivalents for invoking this method: - * - * {{{ - * record.apply(field) - * record(field) - * }}} - * - * @param field the field to access in this record - * @return the value associated with `field`. - * @throws IllegalStateException - */ - @throws(classOf[IllegalStateException]) - def apply[A](field: Field[A]): A = - getOrInitializeEntry(field).value.asInstanceOf[A] - - /** - * Locks the current value for a given `field` in this record, preventing further `update`s. If - * there is no value currently assigned (either explicitly by a previous `update`, or by the - * field's declared default), this will throw an `IllegalStateException`, indicating the field - * was never initialized. - * - * @param field the field to lock in this record - * @return this record - * @throws IllegalStateException - */ - @throws(classOf[IllegalStateException]) - def lock(field: Field[_]): Record = { - getOrInitializeEntry(field).locked = true - this - } - - /** - * Assigns (or reassigns) a `value` to a `field` in this record. If this field was previously - * `lock`ed, this will throw an `IllegalStateException` to indicate failure. - * - * Note that Scala provides two syntactic equivalents for invoking this method: - * - * {{{ - * record.update(field, value) - * record(field) = value - * }}} - * - * @param field the field to assign in this record - * @param value the value to assign to `field` in this record - * @return this record - * @throws IllegalStateException - */ - @throws(classOf[IllegalStateException]) - def update[A](field: Field[A], value: A): Record = { - val entry = fields.get(field) - if (entry eq null) { - fields.put(field, new Entry(value)) - } else if (entry.locked) { - throw new IllegalStateException( - s"attempt to assign $value to a locked field (with current value ${entry.value})" - ) - } else { - entry.value = value - } - this - } - - /** - * Assigns (or reassigns) a `value` to a `field` in this record, and locks it to prevent further - * `update`s. This method is provided for convenience only; the following are equivalent: - * - * {{{ - * record.updateAndLock(field, value) - * record.update(field, value).lock(field) - * }}} - * - * @param field the field to assign and lock in this record - * @param value the value to assign to `field` in this record - * @return this record - * @throws IllegalStateException - */ - @throws(classOf[IllegalStateException]) - def updateAndLock[A](field: Field[A], value: A): Record = - update(field, value).lock(field) - - private[this] def copyFields(): IdentityHashMap[Field[_], Entry] = { - val newFields = new IdentityHashMap[Field[_], Entry] - val iter = fields.entrySet().iterator() - while (iter.hasNext()) { - val kv = iter.next() - val entry = kv.getValue() - val newEntry = new Entry(entry.value) - newEntry.locked = entry.locked - newFields.put(kv.getKey(), newEntry) - } - newFields - } - - /** - * Create a copy of this record. Fields are locked in the copy iff they were locked in the - * original record. - * - * @return a copy of this record - */ - def copy(): Record = { - new Record(copyFields()) - } - - /** - * Create a copy of this record with `value` assigned to `field`. `field` will be locked in the - * copy iff it was present and locked in the original record. If `field` was not present in the - * original then the following are equivalent: - * - * {{{ - * record.copy(field, value) - * record.copy().update(field, value) - * }}} - * - * @param field the field to assign in the copy - * @param value the value to assign to `field` in the copy - * @return a copy of this record - */ - def copy[A](field: Field[A], value: A): Record = { - val newFields = copyFields() - val entry = newFields.get(field) - if (entry eq null) { - newFields.put(field, new Entry(value)) - } else { - entry.value = value - } - new Record(newFields) - } - } - - /** - * Creates a new [[com.twitter.collection.RecordSchema.Record Record]] from this Schema. - * - * @return a new [[com.twitter.collection.RecordSchema.Record Record]] - */ - def newRecord(): Record = new Record - - /** - * Field is a handle used to access some corresponding value in a - * [[com.twitter.collection.RecordSchema.Record Record]]. A field may also declare a default, - * which is computed for each record it is associated with, at most once per record instance. - */ - sealed trait Field[A] { - def default(): A - } - - /** - * Creates a new [[com.twitter.collection.RecordSchema.Field Field]] with no default value, to be - * used only with [[com.twitter.collection.RecordSchema.Record Records]] from this schema. - * - * @return a [[com.twitter.collection.RecordSchema.Field Field]] with no default value - */ - def newField[A](): Field[A] = new Field[A] { - override def default(): A = - throw new IllegalStateException("attempt to access uninitialized field") - } - - /** - * Creates a new [[com.twitter.collection.RecordSchema.Field Field]] with the given - * `defaultSupplier`, to be used only with [[com.twitter.collection.RecordSchema.Record Records]] - * from this schema. - * - * @param defaultSupplier a computation producing the default value to use, when there is no value - * previously assigned to this field in a given record. - * @return a [[com.twitter.collection.RecordSchema.Field Field]] with the given `defaultSupplier` - */ - def newField[A](defaultSupplier: => A): Field[A] = new Field[A] { - override def default(): A = defaultSupplier - } -} diff --git a/util-collection/src/main/scala/com/twitter/util/ImmutableLRU.scala b/util-collection/src/main/scala/com/twitter/util/ImmutableLRU.scala deleted file mode 100644 index 17c04442b7..0000000000 --- a/util-collection/src/main/scala/com/twitter/util/ImmutableLRU.scala +++ /dev/null @@ -1,104 +0,0 @@ -package com.twitter.util - -import scala.collection.SortedMap - -/** - * Immutable implementation of an LRU cache. - */ -object ImmutableLRU { - - /** - * Build an immutable LRU key/value store that cannot grow larger than `maxSize` - */ - def apply[K, V](maxSize: Int): ImmutableLRU[K, V] = { - new ImmutableLRU(maxSize, 0, Map.empty[K, (Long, V)], SortedMap.empty[Long, K]) - } -} - -/** - * An immutable key/value store that evicts the least recently accessed elements - * to stay constrained in a maximum size bound. - */ -// "map" is the backing store used to hold key->(index,value) -// pairs. The index tracks the access time for a particular key. "ord" -// is used to determine the Least-Recently-Used key in "map" by taking -// the minimum index. -class ImmutableLRU[K, V] private ( - maxSize: Int, - idx: Long, - map: Map[K, (Long, V)], - ord: SortedMap[Long, K] -) { - - // Scala's SortedMap requires a key ordering; ImmutableLRU doesn't - // care about pulling a minimum value out of the SortedMap, so the - // following kOrd treats every value as equal. - protected implicit val kOrd = new Ordering[K] { def compare(l: K, r: K) = 0 } - - /** - * the number of entries in the cache - */ - def size: Int = map.size - - /** - * the `Set` of all keys in the LRU - * @note accessing this set does not update the element LRU ordering - */ - def keySet: Set[K] = map.keySet - - /** - * Build a new LRU containing the given key/value - * @return a tuple with of the following two items: - * _1 represents the evicted entry (if the given lru is at the maximum size) - * or None if the lru is not at capacity yet. - * _2 is the new lru with the given key/value pair inserted. - */ - def +(kv: (K, V)): (Option[K], ImmutableLRU[K, V]) = { - val (key, value) = kv - val newIdx = idx + 1 - val newMap = map + (key -> ((newIdx, value))) - // Now update the ordered cache: - val baseOrd = map.get(key).map { case (id, _) => ord - id }.getOrElse(ord) - val ordWithNewKey = baseOrd + (newIdx -> key) - // Do we need to remove an old key: - val (evicts, finalMap, finalOrd) = if (ordWithNewKey.size > maxSize) { - val (minIdx, eKey) = ordWithNewKey.min - (Some(eKey), newMap - eKey, ordWithNewKey - minIdx) - } else { - (None, newMap, ordWithNewKey) - } - (evicts, new ImmutableLRU[K, V](maxSize, newIdx, finalMap, finalOrd)) - } - - /** - * If the key is present in the cache, returns the pair of - * Some(value) and the cache with the key's entry as the most recently accessed. - * Else, returns None and the unmodified cache. - */ - def get(k: K): (Option[V], ImmutableLRU[K, V]) = { - val (optionalValue, lru) = remove(k) - val newLru = optionalValue.map { v => - (lru + (k -> v))._2 - } getOrElse (lru) - (optionalValue, newLru) - } - - /** - * If the key is present in the cache, returns the pair of - * Some(value) and the cache with the key removed. - * Else, returns None and the unmodified cache. - */ - def remove(k: K): (Option[V], ImmutableLRU[K, V]) = - map - .get(k) - .map { - case (kidx, v) => - val newMap = map - k - val newOrd = ord - kidx - // Note we don't increase the idx on a remove, only on put: - (Some(v), new ImmutableLRU[K, V](maxSize, idx, newMap, newOrd)) - } - .getOrElse((None, this)) - - override def toString = { "ImmutableLRU(" + map.toList.mkString(",") + ")" } -} diff --git a/util-collection/src/main/scala/com/twitter/util/LruMap.scala b/util-collection/src/main/scala/com/twitter/util/LruMap.scala deleted file mode 100644 index a0791dc26d..0000000000 --- a/util-collection/src/main/scala/com/twitter/util/LruMap.scala +++ /dev/null @@ -1,70 +0,0 @@ -package com.twitter.util - -import java.{util => ju} -import java.util.LinkedHashMap -import scala.collection.JavaConverters._ -import scala.collection.mutable.{Map, MapLike, SynchronizedMap} - -/** - * A wrapper trait for java.util.Map implementations to make them behave as scala Maps. - * This is useful if you want to have more specifically-typed wrapped objects instead - * of the generic maps returned by JavaConverters - */ -trait JMapWrapperLike[A, B, +Repr <: MapLike[A, B, Repr] with Map[A, B]] - extends Map[A, B] - with MapLike[A, B, Repr] { - def underlying: ju.Map[A, B] - - override def size = underlying.size - - override def get(k: A) = underlying.asScala.get(k) - - override def +=(kv: (A, B)): this.type = { underlying.put(kv._1, kv._2); this } - override def -=(key: A): this.type = { underlying remove key; this } - - override def put(k: A, v: B): Option[B] = underlying.asScala.put(k, v) - - override def update(k: A, v: B): Unit = { underlying.put(k, v) } - - override def remove(k: A): Option[B] = underlying.asScala.remove(k) - - override def clear() = underlying.clear() - - override def empty: Repr = null.asInstanceOf[Repr] - - override def iterator = underlying.asScala.iterator -} - -object LruMap { - - // initial capacity and load factor are the normal defaults for LinkedHashMap - def makeUnderlying[K, V](maxSize: Int): ju.Map[K, V] = - new LinkedHashMap[K, V]( - 16, /* initial capacity */ - 0.75f, /* load factor */ - true /* access order (as opposed to insertion order) */ - ) { - override protected def removeEldestEntry(eldest: ju.Map.Entry[K, V]): Boolean = { - this.size() > maxSize - } - } -} - -/** - * A scala `Map` backed by a [[java.util.LinkedHashMap]] - */ -class LruMap[K, V](val maxSize: Int, val underlying: ju.Map[K, V]) - extends JMapWrapperLike[K, V, LruMap[K, V]] { - override def empty: LruMap[K, V] = new LruMap[K, V](maxSize) - def this(maxSize: Int) = this(maxSize, LruMap.makeUnderlying(maxSize)) -} - -/** - * A synchronized scala `Map` backed by an [[java.util.LinkedHashMap]] - */ -class SynchronizedLruMap[K, V](maxSize: Int, underlying: ju.Map[K, V]) - extends LruMap[K, V](maxSize, ju.Collections.synchronizedMap(underlying)) - with SynchronizedMap[K, V] { - override def empty: SynchronizedLruMap[K, V] = new SynchronizedLruMap[K, V](maxSize) - def this(maxSize: Int) = this(maxSize, LruMap.makeUnderlying(maxSize)) -} diff --git a/util-collection/src/main/scala/com/twitter/util/MapToSetAdapter.scala b/util-collection/src/main/scala/com/twitter/util/MapToSetAdapter.scala deleted file mode 100644 index c2d5da5c68..0000000000 --- a/util-collection/src/main/scala/com/twitter/util/MapToSetAdapter.scala +++ /dev/null @@ -1,17 +0,0 @@ -package com.twitter.util - -import scala.collection.mutable - -class MapToSetAdapter[A](map: mutable.Map[A, A]) extends mutable.Set[A] { - def +=(elem: A): this.type = { - map(elem) = elem - this - } - def -=(elem: A): this.type = { - map -= elem - this - } - override def size: Int = map.size - def iterator: Iterator[A] = map.keysIterator - def contains(elem: A): Boolean = map.contains(elem) -} diff --git a/util-collection/src/test/scala/BUILD b/util-collection/src/test/scala/BUILD deleted file mode 100644 index 92e867cbc5..0000000000 --- a/util-collection/src/test/scala/BUILD +++ /dev/null @@ -1,12 +0,0 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalacheck", - "3rdparty/jvm/org/scalatest", - "util/util-collection/src/main/scala", - "util/util-core/src/main/scala", - ], -) diff --git a/util-collection/src/test/scala/com/twitter/collection/GenerationalQueueTest.scala b/util-collection/src/test/scala/com/twitter/collection/GenerationalQueueTest.scala deleted file mode 100644 index a698bfc506..0000000000 --- a/util-collection/src/test/scala/com/twitter/collection/GenerationalQueueTest.scala +++ /dev/null @@ -1,132 +0,0 @@ -package com.twitter.collection - -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner - -import com.twitter.conversions.time._ -import com.twitter.util.{Duration, Time} - -@RunWith(classOf[JUnitRunner]) -class GenerationalQueueTest extends FunSuite { - - def genericGenerationalQueueTest[A]( - name: String, - newQueue: () => GenerationalQueue[String], - timeout: Duration - ): Unit = { - - test("%s: Don't collect fresh data".format(name)) { - Time.withCurrentTimeFrozen { _ => - val queue = newQueue() - queue.add("foo") - queue.add("bar") - val collectedValue = queue.collect(timeout) - assert(collectedValue.isEmpty) - } - } - - test("%s: Don't collect old data recently refreshed".format(name)) { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = newQueue() - queue.add("foo") - t += timeout * 3 - queue.add("bar") - t += timeout * 3 - queue.touch("foo") - t += timeout * 3 - val collectedValue = queue.collect(timeout) - assert(!collectedValue.isEmpty) - assert(collectedValue.get == "bar") - } - } - - test("%s: collect one old data".format(name)) { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = newQueue() - queue.add("foo1") - queue.add("foo2") - queue.add("foo3") - - t += timeout * 3 - queue.add("bar") - - val collectedData = queue.collect(timeout + 1.millisecond) - assert(!collectedData.isEmpty) - assert(collectedData.get.indexOf("foo") != -1) - } - } - - test("%s: collectAll old data".format(name)) { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = newQueue() - queue.add("foo") - queue.add("bar") - - t += timeout - queue.add("foo2") - queue.add("bar2") - - t += timeout - val collectedValues = queue.collectAll(timeout + 1.millisecond) - assert(collectedValues.size == 2) - assert(collectedValues.toSeq.contains("foo")) - assert(collectedValues.toSeq.contains("bar")) - } - } - - test("%s: generate a long chain of bucket (if applicable)".format(name)) { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = newQueue() - queue.add("foo") - t += timeout * 2 - queue.add("bar") - t += timeout * 2 - queue.add("foobar") - t += timeout * 2 - queue.add("barfoo") - - val collectedValues = queue.collect(timeout + 1.millisecond) - assert(!collectedValues.isEmpty) - assert(List("foo", "bar").contains(collectedValues.get)) - } - } - - test("%s: Don't collect old data".format(name)) { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = newQueue() - t += 3.seconds - queue.add("1") - queue.add("2") - queue.add("3") - t += 4.seconds - assert(!queue.collect(2.seconds).isEmpty) - } - } - } - - val generationalQueue = () => new ExactGenerationalQueue[String]() - genericGenerationalQueueTest("ExactGenerationalQueue", generationalQueue, 100.milliseconds) - - val timeout = 100.milliseconds - val bucketQueue = () => new BucketGenerationalQueue[String](timeout) - genericGenerationalQueueTest("BucketGenerationalQueue", bucketQueue, timeout) - - test("BucketGenerationalQueue check removals") { - var t = Time.now - Time.withTimeFunction(t) { _ => - val queue = bucketQueue() - queue.add("1") - queue.add("2") - queue.remove("1") - queue.remove("2") - t += timeout * 3 - assert(queue.collect(timeout).isEmpty) - } - } -} diff --git a/util-collection/src/test/scala/com/twitter/collection/RecordSchemaTest.scala b/util-collection/src/test/scala/com/twitter/collection/RecordSchemaTest.scala deleted file mode 100644 index 913cf21bc2..0000000000 --- a/util-collection/src/test/scala/com/twitter/collection/RecordSchemaTest.scala +++ /dev/null @@ -1,130 +0,0 @@ -package com.twitter.collection - -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class RecordSchemaTest extends FunSuite { - - val schema = new RecordSchema - val field = schema.newField[Object]() - val fieldWithDefault = schema.newField[Object](new Object) - val fields = Seq(field, fieldWithDefault) - - test("apply should throw IllegalStateException when field is uninitialized") { - val record = schema.newRecord() - intercept[IllegalStateException] { - record(field) - } - } - - test("apply should compute, store and return default when field is initialized with default") { - val record = schema.newRecord() - assert(record(fieldWithDefault) eq record(fieldWithDefault)) - } - - test("apply should return field value when field is explicitly initialized") { - val record = schema.newRecord() - val value = new Object - - for (f <- fields) { - record(f) = value - assert(record(f) eq value) - } - } - - test("lock should throw IllegalStateException when field is uninitialized") { - val record = schema.newRecord() - intercept[IllegalStateException] { - record.lock(field) - } - } - - test("lock should compute and store default when field is initialized with default") { - val record = schema.newRecord() - record.lock(fieldWithDefault) - assert(record(fieldWithDefault) ne null) - } - - test("update should reassign when field is not locked") { - val record = schema.newRecord() - val value = new Object - - for (f <- fields) { - record(f) = new Object - record(f) = value - assert(record(f) eq value) - } - } - - test("update should throw IllegalStateException when field is locked") { - val record = schema.newRecord() - val value = new Object - - for (f <- fields) { - record(f) = value - record.lock(f) - intercept[IllegalStateException] { - record(f) = value - } - } - } - - test("updateAndLock should update and lock") { - val record = schema.newRecord() - val value = new Object - - for (f <- fields) { - record.updateAndLock(f, value) - intercept[IllegalStateException] { - record(f) = value - } - } - } - - test("copy should copy") { - val record = schema.newRecord() - val value = new Object - - for (f <- fields) { - record.update(f, value) - val copy = record.copy() - assert(record(f) eq copy(f)) - } - } - - test("copy should not be modified when the original is updated") { - val record = schema.newRecord() - val copy = record.copy - - for (f <- fields) { - record.update(f, new Object) - val copy = record.copy() - record.update(f, new Object) - assert(record(f) ne copy(f)) - } - } - - test("locked state should be copied") { - val record = schema.newRecord() - - for (f <- fields) { - record.updateAndLock(f, new Object) - intercept[IllegalStateException] { - record.copy().update(f, new Object) - } - } - } - - test("copy should be able to overwrite a locked field") { - val record = schema.newRecord() - - for (f <- fields) { - record.updateAndLock(f, new Object) - val value = new Object - val copy = record.copy(f, value) - assert(copy(f) eq value) - } - } -} diff --git a/util-collection/src/test/scala/com/twitter/util/ImmutableLRUTest.scala b/util-collection/src/test/scala/com/twitter/util/ImmutableLRUTest.scala deleted file mode 100644 index f5a85ee154..0000000000 --- a/util-collection/src/test/scala/com/twitter/util/ImmutableLRUTest.scala +++ /dev/null @@ -1,76 +0,0 @@ -package com.twitter.util - -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalacheck.Gen -import org.scalacheck.Arbitrary.arbitrary -import org.scalatest.prop.GeneratorDrivenPropertyChecks - -@RunWith(classOf[JUnitRunner]) -class ImmutableLRUTest extends FunSuite with GeneratorDrivenPropertyChecks { - - // don't waste too much time testing this and keep things small - implicit override val generatorDrivenConfig = - PropertyCheckConfiguration(minSuccessful = 5, minSize = 2, sizeRange = 8) - - test("ImmutableLRU insertion") { - forAll(Gen.zip(Gen.identifier, arbitrary[Int])) { - case (s: String, i: Int) => - val lru = ImmutableLRU[String, Int](50) - - // first-time insertion should not evict - val (key, lru2) = lru + (s -> i) - assert(key == None) - - // test get method - val (key2, _) = lru2.get(s) - assert(key2 == Some(i)) - } - } - - // given a list of entries, build an LRU containing them - private def buildLRU[V]( - lru: ImmutableLRU[String, V], - entries: List[(String, V)] - ): ImmutableLRU[String, V] = { - entries match { - case Nil => lru - case head :: tail => buildLRU((lru + head)._2, tail) - } - } - - test("ImmutableLRU eviction") { - forAll(LRUEntriesGenerator[Int]) { entries => - val lru = buildLRU(ImmutableLRU[String, Int](4), entries) - assert(lru.keySet == entries.map(_._1).slice(entries.size - 4, entries.size).toSet) - } - } - - test("ImmutableLRU removal") { - val gen = for { - entries <- LRUEntriesGenerator[Double] - entry <- Gen.oneOf(entries) - } yield (entries, entry._1, entry._2) - - forAll(gen) { - case (entries, k, v) => - val lru = buildLRU(ImmutableLRU[String, Double](entries.size + 1), entries) - - // make sure it's there - assert(lru.get(k)._1 == Some(v)) - - // remove once - val (key, newLru) = lru.remove(k) - assert(key == Some(v)) - - // remove it again - val (key2, _) = newLru.remove(k) - assert(key2 == None) - - // still should be able to get it out of the original - val (key3, _) = lru.remove(k) - assert(key3 == Some(v)) - } - } -} diff --git a/util-collection/src/test/scala/com/twitter/util/LRUEntriesGenerator.scala b/util-collection/src/test/scala/com/twitter/util/LRUEntriesGenerator.scala deleted file mode 100644 index fdbdba0399..0000000000 --- a/util-collection/src/test/scala/com/twitter/util/LRUEntriesGenerator.scala +++ /dev/null @@ -1,20 +0,0 @@ -package com.twitter.util - -import org.scalacheck.Arbitrary.arbitrary -import org.scalacheck.{Gen, Arbitrary} - -object LRUEntriesGenerator { - - /** - * Generate a nonempty list of map entries keyed with strings - * and having any arbitrarily typed values. - */ - def apply[V: Arbitrary]: Gen[List[(String, V)]] = { - val gen = for { - size <- Gen.choose(1, 200) - uniqueKeys <- Gen.listOfN(size, Gen.identifier).map(_.toSet) - vals <- Gen.listOfN(uniqueKeys.size, arbitrary[V]) - } yield uniqueKeys.toList.zip(vals) - gen suchThat (_.nonEmpty) - } -} diff --git a/util-collection/src/test/scala/com/twitter/util/LRUMapTest.scala b/util-collection/src/test/scala/com/twitter/util/LRUMapTest.scala deleted file mode 100644 index 2dea99015b..0000000000 --- a/util-collection/src/test/scala/com/twitter/util/LRUMapTest.scala +++ /dev/null @@ -1,50 +0,0 @@ -package com.twitter.util - -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalacheck.Gen -import org.scalatest.prop.GeneratorDrivenPropertyChecks - -@RunWith(classOf[JUnitRunner]) -class LRUMapTest extends FunSuite with GeneratorDrivenPropertyChecks { - - // don't waste too much time testing this and keep things small - implicit override val generatorDrivenConfig = - PropertyCheckConfiguration(minSuccessful = 5, minSize = 2, sizeRange = 8) - - test("LRUMap creation") { - forAll(Gen.choose(1, 200)) { size => - val lru = new LruMap[String, String](size) - assert(lru.maxSize == size) - - val slru = new SynchronizedLruMap[String, String](size) - assert(slru.maxSize == size) - } - } - - test("LRUMap insertion") { - forAll(LRUEntriesGenerator[Int]) { entries => - val lru = new LruMap[String, Int](entries.size + 10) // headroom - for (entry <- entries) { - lru += entry - } - - for ((key, value) <- entries) { - assert(lru.get(key) == Some(value)) - } - } - } - - test("LRUMap eviction") { - forAll(LRUEntriesGenerator[Double]) { entries => - val slru = new SynchronizedLruMap[String, Double](5) - for (entry <- entries) { - slru += entry - } - - val expectedKeys = entries.slice(entries.size - 5, entries.size).map(_._1) - assert(slru.keySet == expectedKeys.toSet) - } - } -} diff --git a/util-core/BUILD b/util-core/BUILD index bc41ada0da..a2edd7cd0d 100644 --- a/util-core/BUILD +++ b/util-core/BUILD @@ -1,14 +1,75 @@ target( + tags = ["bazel-compatible"], dependencies = [ + ":scala", "util/util-core/src/main/java", - "util/util-core/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-core/src/test/java", "util/util-core/src/test/scala", ], ) + +# (DPB-10452) This target definition lives here because it needs to accommodate util-core +# code resolution in IntelliJ as well as `util-core-util` which contains packages +# com.twitter.concurrent, com.twitter.conversions, and com.twitter.util +scala_library( + name = "util-core-util", + sources = [ + "concurrent-extra/src/main/scala/com/twitter/concurrent/*.scala", + "src/main/scala-2.12-/com/twitter/conversions/*.scala", + "src/main/scala-2.12-/com/twitter/util/*.scala", + "src/main/scala/com/twitter/util/*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-core-util", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scala-lang:scala-reflect", + "3rdparty/jvm/org/scala-lang/modules:scala-collection-compat", + "3rdparty/jvm/org/scala-lang/modules:scala-parser-combinators", + "util/util-function/src/main/java", + ], + exports = [ + "3rdparty/jvm/org/scala-lang/modules:scala-collection-compat", + "3rdparty/jvm/org/scala-lang/modules:scala-parser-combinators", + "util/util-function/src/main/java", + ], +) + +scala_library( + # This target exists only to act as an aggregator jar. + name = "scala", + sources = ["src/main/scala/PantsWorkaroundCore.scala"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-core", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + ":util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", + ], + exports = [ + ":util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", + ], +) diff --git a/util-core/OWNERS b/util-core/OWNERS deleted file mode 100644 index 97007381d5..0000000000 --- a/util-core/OWNERS +++ /dev/null @@ -1,7 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta -ryano -vkostyukov diff --git a/util-core/PROJECT b/util-core/PROJECT new file mode 100644 index 0000000000..009bd1c5fb --- /dev/null +++ b/util-core/PROJECT @@ -0,0 +1,5 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan diff --git a/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/ForkingScheduler.scala b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/ForkingScheduler.scala new file mode 100644 index 0000000000..467489178d --- /dev/null +++ b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/ForkingScheduler.scala @@ -0,0 +1,104 @@ +package com.twitter.concurrent + +import com.twitter.util.Future +import java.util.concurrent.Executor +import java.util.concurrent.ExecutorService + +/** + * A scheduler that provides methods to fork the execution of + * `Future` computations and behaves similarly to a thread pool + */ +trait ForkingScheduler extends Scheduler { + + /** + * Forks the execution of a `Future` computation even if + * the scheduler is overloaded. + * + * @param f the Future computation to be forked + * @tparam T the type of the `Future` computation + * @return a `Future` to be satisfied when the forked `Future` + * is completed + */ + def fork[T](f: => Future[T]): Future[T] + + /** + * Tries to fork the execution of a `Future` computation and + * returns empty in case the scheduler is overloaded. + * + * @param f the Future computation to be forked + * @tparam T the type of the `Future` computation + * @return A `Future` with `None` if the scheduler is overloaded + * or `Some[T]` if the task is successfully forked and + * executed. + */ + def tryFork[T](f: => Future[T]): Future[Option[T]] + + /** + * Forks the execution of a `Future` computation using the + * provided executor. + * + * @param executor the executor to be used to run the `Future` + * computation + * @param f the `Future` computation to be forked + * @tparam T the type of the `Future` computation + * @return a `Future` to be satisfied when the forked `Future` + * is completed + */ + def fork[T](executor: Executor)(f: => Future[T]): Future[T] + + /** + * Creates a wrapper of this scheduler with a fixed max + * synchronous concurrency. The concurrency limit is + * applied to synchronous execution so tasks can be + * running in parallel if they're suspended at asynchronous + * boundaries but once they are scheduled for execution + * only `concurrency` tasks can be running in parallel. + * + * The original forking scheduler remains + * unchanged and is used as the underlying implementation + * of the returned scheduler. + * + * @param concurrency the max concurrency + * @param maxWaiters max number of pending tasks. When this limit + * is reached, the scheduler starts to return + * `RejectedExecutionException`s + * @return the wrapper with a fixed max concurrency + */ + def withMaxSyncConcurrency(concurrency: Int, maxWaiters: Int): ForkingScheduler + + /** + * Creates an `ExecutorService` wrapper for this forking + * scheduler. The original forking scheduler remains + * unchanged and is used as the underlying implementation + * of the returned executor. If the executor service is + * shut down, only the wrapper is shut down and the original + * forking scheduler continues operating. + * + * @return the executor service wrapper + */ + def asExecutorService(): ExecutorService + + /** + * Indicates if `FuturePool`s should have their tasks + * redirected to the forking scheduler. + * + * @return if enabled + */ + def redirectFuturePools(): Boolean +} + +object ForkingScheduler { + + /** + * Returns empty if the current scheduler isn't a forking + * scheduler. Returns the forking scheduler otherwise. + * @return an optional of the forking scheduler + */ + def apply(): Option[ForkingScheduler] = + Scheduler() match { + case f: ForkingScheduler => + Some(f) + case _ => + None + } +} diff --git a/util-core/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala similarity index 89% rename from util-core/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala rename to util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala index 8fb6807a5b..d22e7fc462 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala +++ b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/NamedPoolThreadFactory.scala @@ -4,10 +4,10 @@ import java.util.concurrent.ThreadFactory import java.util.concurrent.atomic.AtomicInteger /** - * A [[java.util.concurrent.ThreadFactory]] which creates threads with a name + * A `java.util.concurrent.ThreadFactory` which creates threads with a name * indicating the pool from which they originated. * - * A new [[java.lang.ThreadGroup]] (named `name`) is created as a sub-group of + * A new `java.lang.ThreadGroup` (named `name`) is created as a sub-group of * whichever group to which the thread that created the factory belongs. Each * thread created by this factory will be a member of this group and have a * unique name including the group name and an monotonically increasing number. diff --git a/util-core/src/main/scala/com/twitter/concurrent/Offer.scala b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Offer.scala similarity index 95% rename from util-core/src/main/scala/com/twitter/concurrent/Offer.scala rename to util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Offer.scala index 97d324559d..4556a74cf1 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Offer.scala +++ b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Offer.scala @@ -2,6 +2,8 @@ package com.twitter.concurrent import com.twitter.util.{Await, Duration, Future, Promise, Time, Timer} import com.twitter.util.Return +import scala.annotation.varargs +import scala.collection.compat.immutable.ArraySeq import scala.util.Random /** @@ -91,9 +93,7 @@ trait Offer[+T] { self => /** * Like {{map}}, but to a constant (call-by-name). */ - def const[U](f: => U): Offer[U] = map { _ => - f - } + def const[U](f: => U): Offer[U] = map { _ => f } /** * Java-friendly analog of `const()`. @@ -119,9 +119,7 @@ trait Offer[+T] { self => val ourTx = self.prepare() if (ourTx.isDefined) ourTx else { - ourTx.foreach { tx => - tx.nack() - } + ourTx.foreach { tx => tx.nack() } ourTx.raise(LostSynchronization) other.prepare() } @@ -187,7 +185,7 @@ object Offer { * A constant offer: synchronizes the given value always. This is * call-by-name and a new value is produced for each `prepare()`. * - * Note: Updates here must also be done at [[com.twitter.concurrent.Offers.newConstOffer()]]. + * Note: Updates here must also be done at `Offers.newConstOffer`. */ def const[T](x: => T): Offer[T] = new Offer[T] { def prepare(): Future[Tx[T]] = Future.value(Tx.const(x)) @@ -206,6 +204,7 @@ object Offer { * The offer that chooses exactly one of the given offers. If there are any * Offers that are synchronizable immediately, one is chosen at random. */ + @varargs def choose[T](evs: Offer[T]*): Offer[T] = choose(rng, evs) /** @@ -279,7 +278,7 @@ object Offer { if (foundPos >= 0) { updateLosers(foundPos, prepd) } else { - Future.selectIndex(prepd) flatMap { winPos => + Future.selectIndex(ArraySeq.unsafeWrapArray(prepd)) flatMap { winPos => updateLosers(winPos, prepd) } } @@ -312,6 +311,6 @@ object Offer { } object LostSynchronization extends Exception { - override def fillInStackTrace = this + override def fillInStackTrace: LostSynchronization.type = this } } diff --git a/util-core/src/main/scala/com/twitter/concurrent/Scheduler.scala b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Scheduler.scala similarity index 85% rename from util-core/src/main/scala/com/twitter/concurrent/Scheduler.scala rename to util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Scheduler.scala index 1f21c3366a..0927ff875f 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Scheduler.scala +++ b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Scheduler.scala @@ -4,11 +4,12 @@ import com.twitter.util.Awaitable.CanAwait import java.util.ArrayDeque import java.util.concurrent._ import java.util.concurrent.atomic.AtomicInteger -import java.util.logging.{Level, Logger} +import java.util.logging.Level +import java.util.logging.Logger import scala.collection.mutable /** - * An interface for scheduling [[java.lang.Runnable]] tasks. + * An interface for scheduling `java.lang.Runnable` tasks. */ trait Scheduler { @@ -23,37 +24,6 @@ trait Scheduler { */ def flush(): Unit - // A note on Hotspot's ThreadMXBean's CPU time. On Linux, this - // uses clock_gettime[1] which should both be fast and accurate. - // - // On OSX, the Mach thread_info call is used. - // - // [1] http://linux.die.net/man/3/clock_gettime - // - // Note, these clock stats were decommissioned as they only make - // sense relative to `wallTime` and the tracking error we have - // experienced `wallTime` and `*Time` make it impossible to - // use these reliably. It is not worth the performance and - // code complexity to support them. - - /** - * The amount of User time that's been scheduled as per ThreadMXBean. - */ - @deprecated("schedulers no longer export this", "2015-01-10") - def usrTime: Long = -1L - - /** - * The amount of CPU time that's been scheduled as per ThreadMXBean. - */ - @deprecated("schedulers no longer export this", "2015-01-10") - def cpuTime: Long = -1L - - /** - * Total walltime spent in the scheduler. - */ - @deprecated("schedulers no longer export this", "2015-01-10") - def wallTime: Long = -1L - /** * The number of dispatches performed by this scheduler. */ @@ -63,7 +33,7 @@ trait Scheduler { * Total time spent doing [[blocking]] operations, in nanoseconds. * * This should only include time spent on threads where - * [[CanAwait.trackElapsedBlocking]] returns `true`. + * [[com.twitter.util.Awaitable.CanAwait.trackElapsedBlocking]] returns `true`. * * @return -1 if the [[Scheduler]] does not support tracking this. * @@ -138,8 +108,8 @@ private[concurrent] object LocalScheduler { private[this] val rs = new ArrayDeque[Runnable] private[this] var running = false - @volatile var numDispatches = 0L - @volatile var blockingNanos = 0L + @volatile var numDispatches: Long = 0L + @volatile var blockingNanos: Long = 0L override def toString: String = s"Activation(${System.identityHashCode(this)})" @@ -233,7 +203,7 @@ class LocalScheduler(lifo: Boolean) extends Scheduler { // use weak refs to prevent Activations from causing a memory leak // thread-safety provided by synchronizing on `activations` - private[this] val activations = new mutable.WeakHashMap[Activation, Boolean]() + private[this] val activations = new mutable.WeakHashMap[Activation, java.lang.Boolean]() private[this] val local = new ThreadLocal[Activation]() @@ -289,7 +259,7 @@ class LocalScheduler(lifo: Boolean) extends Scheduler { /** * A named Scheduler mix-in that causes submitted tasks to be dispatched according to - * an [[java.util.concurrent.ExecutorService]] created by an abstract factory + * an `java.util.concurrent.ExecutorService` created by an abstract factory * function. */ trait ExecutorScheduler { self: Scheduler => @@ -335,10 +305,8 @@ trait ExecutorScheduler { self: Scheduler => * A scheduler that dispatches directly to an underlying Java * cached threadpool executor. */ -class ThreadPoolScheduler( - val name: String, - val executorFactory: ThreadFactory => ExecutorService -) extends Scheduler +class ThreadPoolScheduler(val name: String, val executorFactory: ThreadFactory => ExecutorService) + extends Scheduler with ExecutorScheduler { def this(name: String) = this(name, Executors.newCachedThreadPool(_)) } @@ -354,8 +322,8 @@ class ThreadPoolScheduler( */ class BridgedThreadPoolScheduler( val name: String, - val executorFactory: ThreadFactory => ExecutorService -) extends Scheduler + val executorFactory: ThreadFactory => ExecutorService) + extends Scheduler with ExecutorScheduler { private[this] val local = new LocalScheduler diff --git a/util-core/src/main/scala/com/twitter/concurrent/Tx.scala b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Tx.scala similarity index 92% rename from util-core/src/main/scala/com/twitter/concurrent/Tx.scala rename to util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Tx.scala index 9cc36fc74e..9e11de9fc9 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Tx.scala +++ b/util-core/concurrent-extra/src/main/scala/com/twitter/concurrent/Tx.scala @@ -50,7 +50,7 @@ object Tx { /** * A `Tx` that will always commit the given value immediately. * - * Note: Updates here must also be done at [[com.twitter.concurrent.Txs.newConstTx()]]. + * Note: Updates here must also be done at `Txs.newConstTx`. */ def const[T](msg: T): Tx[T] = new Tx[T] { def ack() = Future.value(Commit(msg)) @@ -65,7 +65,7 @@ object Tx { /** * A constant `Tx` with the value of `Unit`. */ - val Unit = const(()) + val Unit: Tx[Unit] = const(()) object AlreadyDone extends Exception("Tx is already done") object AlreadyAckd extends Exception("Tx was already ackd") @@ -91,10 +91,12 @@ object Tx { state match { case Idle => val p = new Promise[Result[U]] - state = Ackd(this, { - case true => p.setValue(Commit(msg)) - case false => p.setValue(Abort) - }) + state = Ackd( + this, + { + case true => p.setValue(Commit(msg)) + case false => p.setValue(Abort) + }) p case Ackd(who, confirm) if who ne this => diff --git a/util-core/src/benchmark/scala/BUILD.bazel b/util-core/src/benchmark/scala/BUILD.bazel new file mode 100644 index 0000000000..7f9fa00b5a --- /dev/null +++ b/util-core/src/benchmark/scala/BUILD.bazel @@ -0,0 +1,29 @@ +# Run locally: +# ./bazel run util/util-core/src/benchmark/scala $BENCHMARK_NAME_REGEXP +# +# Run on a remote host: +# Bundle: +# ./bazel bundle util/util-core/src/benchmark/scala +# Copy jars to a target host: +# rsync -Lrv --delete dist/util-core-jmh-bundle* $TARGET_HOST: +# Run on the remote host: +# ssh $TARGET_HOST /usr/lib/jvm/java-11-twitter/bin/java -jar util-core-jmh-bundle/util-core-jmh.jar $BENCHMARK_NAME_REGEXP +jvm_binary( + basename = "util-core-jmh", + main = "org.openjdk.jmh.Main", + runtime_platform = "java11", + dependencies = [ + "util/util-core/src/benchmark/scala:jmh_compiled_benchmark_lib", + ], +) + +scala_benchmark_jmh( + name = "jmh", + sources = [ + "**/*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + dependencies = [ + "util/util-core", + ], +) diff --git a/util-core/src/benchmark/scala/com/twitter/util/LocalBenchmarks.scala b/util-core/src/benchmark/scala/com/twitter/util/LocalBenchmarks.scala new file mode 100644 index 0000000000..0aff5b3e7c --- /dev/null +++ b/util-core/src/benchmark/scala/com/twitter/util/LocalBenchmarks.scala @@ -0,0 +1,29 @@ +package com.twitter.util + +import com.twitter.util.Local.Context +import java.util.concurrent.ThreadLocalRandom +import org.openjdk.jmh.annotations.Benchmark +import org.openjdk.jmh.annotations.Fork +import org.openjdk.jmh.annotations.Scope +import org.openjdk.jmh.annotations.State +import org.openjdk.jmh.annotations.Warmup +import org.openjdk.jmh.infra.Blackhole + +@State(Scope.Benchmark) +@Warmup(iterations = 3) +@Fork(2) +class LocalBenchmarks { + + private[this] val key = new Local.Key + + @Benchmark + def let(blackhole: Blackhole): Unit = { + val randLong = ThreadLocalRandom.current().nextLong() + val newContext = Context.empty.set(key, Some(randLong)) + + Local.let(newContext) { + blackhole.consume(randLong) + } + } + +} diff --git a/util-core/src/main/java/BUILD b/util-core/src/main/java/BUILD index ec624fbc75..4c75a6e28e 100644 --- a/util-core/src/main/java/BUILD +++ b/util-core/src/main/java/BUILD @@ -1,19 +1,10 @@ -java_library( - sources = rglobs("*.java"), - fatal_warnings = True, - provides = artifact( - org = "com.twitter", - name = "util-core-java", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-core/src/main/scala", - "util/util-function/src/main/java", - ], - exports = [ - "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-core/src/main/scala", - "util/util-function/src/main/java", + "util/util-core/src/main/java/com/twitter/concurrent", + "util/util-core/src/main/java/com/twitter/io", + "util/util-core/src/main/java/com/twitter/service", + "util/util-core/src/main/java/com/twitter/util", + "util/util-core/src/main/java/com/twitter/util/javainterop", ], ) diff --git a/util-core/src/main/java/com/twitter/concurrent/BUILD b/util-core/src/main/java/com/twitter/concurrent/BUILD new file mode 100644 index 0000000000..51c4ae41b6 --- /dev/null +++ b/util-core/src/main/java/com/twitter/concurrent/BUILD @@ -0,0 +1,19 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-core-java-concurrent", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], + exports = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], +) diff --git a/util-core/src/main/java/com/twitter/concurrent/Offers.java b/util-core/src/main/java/com/twitter/concurrent/Offers.java index a26569b3a1..8270e6cbb9 100644 --- a/util-core/src/main/java/com/twitter/concurrent/Offers.java +++ b/util-core/src/main/java/com/twitter/concurrent/Offers.java @@ -4,7 +4,7 @@ import com.twitter.util.Function0; import com.twitter.util.Future; import com.twitter.util.Timer; -import scala.collection.JavaConversions; +import scala.collection.JavaConverters; import scala.runtime.BoxedUnit; import java.util.ArrayList; @@ -60,7 +60,8 @@ public static Offer newTimeoutOffer(Duration duration, Timer timer) { * @see Offer$#choose(scala.collection.Seq) */ public static Offer choose(Collection> offers) { - return Offer$.MODULE$.choose(JavaConversions.asScalaBuffer(new ArrayList>(offers))); + scala.collection.immutable.Seq> scalaSeq = JavaConverters.collectionAsScalaIterableConverter(offers).asScala().toVector(); + return Offer$.MODULE$.choose(scalaSeq); } /** diff --git a/util-core/src/main/java/com/twitter/concurrent/Spools.java b/util-core/src/main/java/com/twitter/concurrent/Spools.java index 6343b997b1..b8a34ab030 100644 --- a/util-core/src/main/java/com/twitter/concurrent/Spools.java +++ b/util-core/src/main/java/com/twitter/concurrent/Spools.java @@ -1,7 +1,11 @@ package com.twitter.concurrent; -import scala.collection.JavaConversions; +import com.twitter.util.Future; + +import scala.collection.JavaConverters; import scala.collection.Seq; +import scala.collection.immutable.List; +import scala.collection.immutable.Nil$; import java.util.ArrayList; import java.util.Collection; @@ -21,9 +25,18 @@ private Spools() { } /** * Creates a new `Spool` of given `elems`. */ + @SuppressWarnings("unchecked") public static Spool newSpool(Collection elems) { - Seq seq = JavaConversions.asScalaBuffer(new ArrayList(elems)).toSeq(); - return new Spool.ToSpool(seq).toSpool(); + List buffer = (List)Nil$.MODULE$; + for (T item : elems) { + buffer = buffer.$colon$colon(item); + } + Spool result = (Spool)EMPTY; + while(!buffer.isEmpty()){ + result = new Spool.Cons(buffer.head(), Future.value(result)); + buffer = (List)buffer.tail(); + } + return result; } /** diff --git a/util-core/src/main/java/com/twitter/io/BUILD b/util-core/src/main/java/com/twitter/io/BUILD new file mode 100644 index 0000000000..febf1f1097 --- /dev/null +++ b/util-core/src/main/java/com/twitter/io/BUILD @@ -0,0 +1,19 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-core-java-io", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], + exports = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], +) diff --git a/util-core/src/main/java/com/twitter/io/Bufs.java b/util-core/src/main/java/com/twitter/io/Bufs.java index 12997d51f5..5ca30866eb 100644 --- a/util-core/src/main/java/com/twitter/io/Bufs.java +++ b/util-core/src/main/java/com/twitter/io/Bufs.java @@ -4,7 +4,10 @@ /** * A Java adaptation of the {@link com.twitter.io.Buf} companion object. + * + * @deprecated This will no longer be necessary when 2.11 support is dropped. 2019-10-04 */ +@Deprecated public final class Bufs { private Bufs() { } diff --git a/util-core/src/main/java/com/twitter/io/Readers.java b/util-core/src/main/java/com/twitter/io/Readers.java deleted file mode 100644 index 7c54857478..0000000000 --- a/util-core/src/main/java/com/twitter/io/Readers.java +++ /dev/null @@ -1,99 +0,0 @@ -/* Copyright 2015 Twitter, Inc. */ -package com.twitter.io; - -import java.io.File; -import java.io.FileNotFoundException; -import java.io.InputStream; -import com.twitter.concurrent.AsyncStream; -import com.twitter.util.Future; -import scala.runtime.BoxedUnit; - -/** - * Java APIs for Reader. - * - * @see com.twitter.io.Reader - */ -public final class Readers { - - private Readers() { throw new IllegalStateException(); } - - /** - * See {@code com.twitter.io.Reader.Null}. - */ - public static final Reader NULL = Reader$.MODULE$.Null(); - - public static Reader newEmptyReader() { - return Reader$.MODULE$.empty(); - } - - /** - * See {@code com.twitter.io.Reader.fromBuf}. - */ - public static Reader newBufReader(Buf buf) { - return Reader$.MODULE$.fromBuf(buf); - } - - /** - * See {@code com.twitter.io.Reader.readAll}. - */ - public static Future readAll(Reader r) { - return Reader$.MODULE$.readAll(r); - } - - /** - * See {@code com.twitter.io.Reader.concat}. - */ - public static Reader concat(AsyncStream> readers) { - return Reader$.MODULE$.concat(readers); - } - - /** - * See {@code com.twitter.io.Reader.copy}. - */ - public static Future copy(Reader r, Writer w) { - return Reader$.MODULE$.copy(r, w); - } - - /** - * See {@code com.twitter.io.Reader.copy}. - */ - public static Future copy(Reader r, Writer w, int readSize) { - return Reader$.MODULE$.copy(r, w, readSize); - } - - /** - * See {@code com.twitter.io.Reader.copyMany}. - */ - public static Future copyMany(AsyncStream> readers, Writer w) { - return Reader$.MODULE$.copyMany(readers, w); - } - - /** - * See {@code com.twitter.io.Reader.copyMany}. - */ - public static Future copyMany(AsyncStream> readers, Writer w, int readSize) { - return Reader$.MODULE$.copyMany(readers, w, readSize); - } - - /** - * See {@code com.twitter.io.Pipe}. - */ - public static Pipe writable() { - return new Pipe<>(); - } - - /** - * See {@code com.twitter.io.Reader.fromFile()}. - */ - public static Reader newFileReader(File f) throws FileNotFoundException { - return Reader$.MODULE$.fromFile(f); - } - - /** - * See {@code com.twitter.io.Reader.fromStream()}. - */ - public static Reader newInputStreamReader(InputStream in) { - return Reader$.MODULE$.fromStream(in); - } - -} diff --git a/util-core/src/main/java/com/twitter/io/Writers.java b/util-core/src/main/java/com/twitter/io/Writers.java deleted file mode 100644 index b89a2afc9a..0000000000 --- a/util-core/src/main/java/com/twitter/io/Writers.java +++ /dev/null @@ -1,29 +0,0 @@ -/* Copyright 2015 Twitter, Inc. */ -package com.twitter.io; - -import java.io.OutputStream; - -/** - * Java APIs for Writer. - * - * @see com.twitter.io.Writer - */ -public final class Writers { - - private Writers() { throw new IllegalStateException(); } - - /** - * See {@code com.twitter.io.Writer.fromOutputStream} - */ - public static Writer.ClosableWriter newOutputStreamWriter(OutputStream out) { - return Writer$.MODULE$.fromOutputStream(out); - } - - /** - * See {@code com.twitter.io.Writer.fromOutputStream} - */ - public static Writer.ClosableWriter newOutputStreamWriter(OutputStream out, int bufsize) { - return Writer$.MODULE$.fromOutputStream(out, bufsize); - } - -} diff --git a/util-core/src/main/java/com/twitter/service/BUILD b/util-core/src/main/java/com/twitter/service/BUILD new file mode 100644 index 0000000000..59f3eda853 --- /dev/null +++ b/util-core/src/main/java/com/twitter/service/BUILD @@ -0,0 +1,12 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-core-java-service", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-core/src/main/java/com/twitter/service/VirtualMachineErrorTerminator.java b/util-core/src/main/java/com/twitter/service/VirtualMachineErrorTerminator.java index c6910d53e4..5f7fd9e649 100644 --- a/util-core/src/main/java/com/twitter/service/VirtualMachineErrorTerminator.java +++ b/util-core/src/main/java/com/twitter/service/VirtualMachineErrorTerminator.java @@ -5,8 +5,8 @@ /** * Utility class to terminate the JVM in case of a virtual machine error. * Classes wishing to use should initialize it early using - * {@link #initialize()}, as it might be unable to initialize properly after - * the virtual machine error has been thrown. See {@link #checkTerminating()} + * `initialize`, as it might be unable to initialize properly after + * the virtual machine error has been thrown. See `checkTerminating` * for the typical usage pattern. * @author Attila Szegedi */ diff --git a/util-core/src/main/java/com/twitter/util/BUILD b/util-core/src/main/java/com/twitter/util/BUILD new file mode 100644 index 0000000000..38f12c51e1 --- /dev/null +++ b/util-core/src/main/java/com/twitter/util/BUILD @@ -0,0 +1,19 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-core-java", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], + exports = [ + "3rdparty/jvm/org/scala-lang:scala-library", + "util/util-core:scala", + ], +) diff --git a/util-core/src/main/java/com/twitter/util/Closables.java b/util-core/src/main/java/com/twitter/util/Closables.java index 5b75507731..f1bd814c54 100644 --- a/util-core/src/main/java/com/twitter/util/Closables.java +++ b/util-core/src/main/java/com/twitter/util/Closables.java @@ -1,6 +1,6 @@ package com.twitter.util; -import scala.collection.JavaConversions; +import scala.collection.JavaConverters; import scala.runtime.BoxedUnit; import java.util.Arrays; @@ -43,18 +43,14 @@ public static Future close(Object o, Duration after) { * @see com.twitter.util.Closable$#all(scala.collection.Seq) */ public static Closable all(Closable... closables) { - return Closable$.MODULE$.all( - JavaConversions.asScalaBuffer(Arrays.asList(closables)) - ); + return Closable$.MODULE$.all(closables); } /** * @see com.twitter.util.Closable$#sequence(scala.collection.Seq) */ public static Closable sequence(Closable... closables) { - return Closable$.MODULE$.sequence( - JavaConversions.asScalaBuffer(Arrays.asList(closables)) - ); + return Closable$.MODULE$.sequence(closables); } /** diff --git a/util-core/src/main/java/com/twitter/util/Monitors.java b/util-core/src/main/java/com/twitter/util/Monitors.java index d8887d75c2..8022586fdc 100644 --- a/util-core/src/main/java/com/twitter/util/Monitors.java +++ b/util-core/src/main/java/com/twitter/util/Monitors.java @@ -14,56 +14,53 @@ public final class Monitors { private Monitors() { } - private static final Monitor$ it = - Monitor$.MODULE$; - /** * The singleton `Monitor` instance that is the Monitor companion object. */ public static Monitor instance() { - return it; + return Monitor$.MODULE$; } /** * Equivalent to calling `Monitor.get`. */ public static Monitor get() { - return it.get(); + return Monitor$.MODULE$.get(); } /** * Equivalent to calling `Monitor.set`. */ public static void set(Monitor monitor) { - it.set(monitor); + Monitor$.MODULE$.set(monitor); } /** * Equivalent to calling `Monitor.using`. */ public static T using(Monitor monitor, Supplier supplier) { - return it.using(monitor, func0(supplier)); + return Monitor$.MODULE$.using(monitor, func0(supplier)); } /** * Equivalent to calling `Monitor.restoring`. */ public static T restoring(Supplier supplier) { - return it.restoring(func0(supplier)); + return Monitor$.MODULE$.restoring(func0(supplier)); } /** * Equivalent to `Monitor.catcher`. */ public static PartialFunction catcher() { - return it.catcher(); + return Monitor$.MODULE$.catcher(); } /** * Equivalent to calling `Monitor.isActive`. */ public static boolean isActive() { - return it.isActive(); + return Monitor$.MODULE$.isActive(); } } diff --git a/util-core/src/main/java/com/twitter/util/Vars.java b/util-core/src/main/java/com/twitter/util/Vars.java index 928c18d1b9..4851666eaf 100644 --- a/util-core/src/main/java/com/twitter/util/Vars.java +++ b/util-core/src/main/java/com/twitter/util/Vars.java @@ -34,8 +34,8 @@ public static Var newVar(T init, Event event) { /** * @see com.twitter.util.Var$#value(Object) */ - public static Var newConstVar(T constant) { - return Var$.MODULE$.value(constant); + public static ConstVar newConstVar(T constant) { + return new ConstVar(constant); } /** diff --git a/util-core/src/main/java/com/twitter/util/javainterop/BUILD b/util-core/src/main/java/com/twitter/util/javainterop/BUILD new file mode 100644 index 0000000000..9ac9a13322 --- /dev/null +++ b/util-core/src/main/java/com/twitter/util/javainterop/BUILD @@ -0,0 +1,17 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-core-java-javainterop", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scala-lang:scala-library", + ], + exports = [ + "3rdparty/jvm/org/scala-lang:scala-library", + ], +) diff --git a/util-core/src/main/java/com/twitter/util/javainterop/Scala.java b/util-core/src/main/java/com/twitter/util/javainterop/Scala.java index 50528ff78f..6990225275 100644 --- a/util-core/src/main/java/com/twitter/util/javainterop/Scala.java +++ b/util-core/src/main/java/com/twitter/util/javainterop/Scala.java @@ -3,14 +3,13 @@ import scala.Option; import scala.Tuple2; -import static scala.collection.JavaConversions.asScalaBuffer; -import static scala.collection.JavaConversions.asScalaSet; +import scala.collection.JavaConverters; /** * A collection of conversions from Java collections to their immutable * Scala variants. * - * See scala.collection.JavaConversions if you do not need immutable + * See scala.collection.JavaConverters if you do not need immutable * collections. */ public final class Scala { @@ -20,7 +19,7 @@ public final class Scala { /** * Converts a {@link java.util.List} to an immutable Scala Seq. * - * See scala.collection.JavaConversions.asScalaBuffer if you do + * See scala.collection.JavaConverters.asScalaBuffer if you do * not need the returned Seq to be immutable. * * @return an empty Seq if the input is null. @@ -32,14 +31,14 @@ public static scala.collection.immutable.Seq asImmutableSeq( if (jList == null) { return scala.collection.immutable.Seq$.MODULE$.empty(); } else { - return asScalaBuffer(jList).toList(); + return JavaConverters.asScalaBuffer(jList).toList(); } } /** * Converts a {@link java.util.Set} to an immutable Scala Set. * - * See scala.collection.JavaConversions.asScalaSet if you do + * See scala.collection.JavaConverters.asScalaSet if you do * not need the returned Set to be immutable. * * @return an empty Set if the input is null. @@ -51,14 +50,14 @@ public static scala.collection.immutable.Set asImmutableSet( if (jSet == null) { return scala.collection.immutable.Set$.MODULE$.empty(); } else { - return asScalaSet(jSet).toSet(); + return JavaConverters.asScalaSet(jSet).toSet(); } } /** * Converts a {@link java.util.Map} to an immutable Scala Map. * - * See scala.collection.JavaConversions.asScalaMap if you do + * See scala.collection.JavaConverters.asScalaMap if you do * not need the returned Map to be immutable. * * @return an empty Map if the input is null. diff --git a/util-core/src/main/scala-2.12-/com/twitter/conversions/SeqUtil.scala b/util-core/src/main/scala-2.12-/com/twitter/conversions/SeqUtil.scala new file mode 100644 index 0000000000..41765424de --- /dev/null +++ b/util-core/src/main/scala-2.12-/com/twitter/conversions/SeqUtil.scala @@ -0,0 +1,6 @@ +package com.twitter.conversions + +object SeqUtil { + def hasKnownSize[A](self: Seq[A]): Boolean = + self.hasDefiniteSize +} diff --git a/util-core/src/main/scala-2.12-/com/twitter/util/HealthyQueue.scala b/util-core/src/main/scala-2.12-/com/twitter/util/HealthyQueue.scala new file mode 100644 index 0000000000..04e56a656f --- /dev/null +++ b/util-core/src/main/scala-2.12-/com/twitter/util/HealthyQueue.scala @@ -0,0 +1,40 @@ +package com.twitter.util + +import scala.collection.mutable + +private[util] class HealthyQueue[A]( + makeItem: () => Future[A], + numItems: Int, + isHealthy: A => Boolean) + extends mutable.Queue[Future[A]] { + + 0.until(numItems) foreach { _ => this += makeItem() } + + override def +=(elem: Future[A]): HealthyQueue.this.type = synchronized { + super.+=(elem) + } + + override def +=:(elem: Future[A]): HealthyQueue.this.type = synchronized { + super.+=:(elem) + } + + override def enqueue(elems: Future[A]*): Unit = synchronized { + super.enqueue(elems: _*) + } + + override def dequeue(): Future[A] = synchronized { + if (isEmpty) throw new NoSuchElementException("queue empty") + + super.dequeue() flatMap { item => + if (isHealthy(item)) { + Future(item) + } else { + val item = makeItem() + synchronized { + super.enqueue(item) + dequeue() + } + } + } + } +} diff --git a/util-core/src/main/scala-2.12-/com/twitter/util/MapUtil.scala b/util-core/src/main/scala-2.12-/com/twitter/util/MapUtil.scala new file mode 100644 index 0000000000..2c9bc8fd9d --- /dev/null +++ b/util-core/src/main/scala-2.12-/com/twitter/util/MapUtil.scala @@ -0,0 +1,20 @@ +package com.twitter.util + +import java.lang.Integer.numberOfLeadingZeros +import scala.collection.mutable + +object MapUtil { + + def newHashMap[K, V](expectedSize: Int, loadFactor: Double = 0.75): mutable.HashMap[K, V] = { + new mutable.HashMap[K, V]() { + this._loadFactor = (loadFactor * 1000).toInt + override protected def initialSize: Int = { + roundUpToNextPowerOfTwo((expectedSize / loadFactor).toInt max 4) + } + } + } + + private[this] def roundUpToNextPowerOfTwo(target: Int): Int = { + 1 << -numberOfLeadingZeros(target - 1) + } +} diff --git a/util-core/src/main/scala-2.13+/com/twitter/conversions/SeqUtil.scala b/util-core/src/main/scala-2.13+/com/twitter/conversions/SeqUtil.scala new file mode 100644 index 0000000000..a9675dd825 --- /dev/null +++ b/util-core/src/main/scala-2.13+/com/twitter/conversions/SeqUtil.scala @@ -0,0 +1,6 @@ +package com.twitter.conversions + +object SeqUtil { + def hasKnownSize[A](self: Seq[A]): Boolean = + self.knownSize != -1 +} diff --git a/util-core/src/main/scala-2.13+/com/twitter/util/HealthyQueue.scala b/util-core/src/main/scala-2.13+/com/twitter/util/HealthyQueue.scala new file mode 100644 index 0000000000..435b94d123 --- /dev/null +++ b/util-core/src/main/scala-2.13+/com/twitter/util/HealthyQueue.scala @@ -0,0 +1,43 @@ +package com.twitter.util + +import scala.collection.mutable + +private[util] class HealthyQueue[A]( + makeItem: () => Future[A], + numItems: Int, + isHealthy: A => Boolean) + extends scala.collection.mutable.Queue[Future[A]] { + + 0.until(numItems) foreach { _ => this += makeItem() } + + override def addOne(elem: Future[A]): HealthyQueue.this.type = synchronized { + super.addOne(elem) + } + + override def prepend(elem: Future[A]): HealthyQueue.this.type = synchronized { + super.prepend(elem) + } + + override def enqueue(elem: Future[A]) = synchronized { + super.enqueue(elem) + } + override def enqueue(elem1: Future[A], elem2: Future[A], elems: Future[A]*) = synchronized { + super.enqueue(elem1, elem2, elems: _*) + } + + override def dequeue(): Future[A] = synchronized { + if (isEmpty) throw new NoSuchElementException("queue empty") + + super.dequeue() flatMap { item => + if (isHealthy(item)) { + Future(item) + } else { + val item = makeItem() + synchronized { + super.enqueue(item) + dequeue() + } + } + } + } +} diff --git a/util-core/src/main/scala-2.13+/com/twitter/util/MapUtil.scala b/util-core/src/main/scala-2.13+/com/twitter/util/MapUtil.scala new file mode 100644 index 0000000000..af84a676ab --- /dev/null +++ b/util-core/src/main/scala-2.13+/com/twitter/util/MapUtil.scala @@ -0,0 +1,10 @@ +package com.twitter.util + +import scala.collection.mutable + +object MapUtil { + + def newHashMap[K, V](expectedSize: Int, loadFactor: Double = 0.75): mutable.HashMap[K, V] = { + new mutable.HashMap[K, V]((expectedSize / loadFactor).toInt, loadFactor) + } +} diff --git a/util-core/src/main/scala/BUILD b/util-core/src/main/scala/BUILD deleted file mode 100644 index d89ef3c23b..0000000000 --- a/util-core/src/main/scala/BUILD +++ /dev/null @@ -1,18 +0,0 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-core", - repo = artifactory, - ), - dependencies = [ - "3rdparty/jvm/org/scala-lang:scala-reflect", - "3rdparty/jvm/org/scala-lang/modules:scala-parser-combinators", - "util/util-function/src/main/java", - ], - exports = [ - "3rdparty/jvm/org/scala-lang/modules:scala-parser-combinators", - "util/util-function/src/main/java", - ], -) diff --git a/util-core/src/main/scala/PantsWorkaroundCore.scala b/util-core/src/main/scala/PantsWorkaroundCore.scala new file mode 100644 index 0000000000..8aa6cef74e --- /dev/null +++ b/util-core/src/main/scala/PantsWorkaroundCore.scala @@ -0,0 +1,4 @@ +// workaround for https://github.com/pantsbuild/pants/issues/6678 +private final class PantsWorkaroundCore { + throw new IllegalStateException("workaround for https://github.com/pantsbuild/pants/issues/6678") +} diff --git a/util-core/src/main/scala/com/twitter/concurrent/AsyncMeter.scala b/util-core/src/main/scala/com/twitter/concurrent/AsyncMeter.scala index b01ee91873..59e61c4d24 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/AsyncMeter.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/AsyncMeter.scala @@ -1,6 +1,5 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ import com.twitter.util._ import java.util.concurrent.{ ArrayBlockingQueue, @@ -25,16 +24,30 @@ private[concurrent] object Period { } object AsyncMeter { - private[concurrent] val MinimumInterval = 1.millisecond + private[concurrent] val MinimumInterval: Duration = Duration.fromMilliseconds(1) /** * Creates an [[AsyncMeter]] that allows smoothed out `permits` per * second, and has a maximum burst size of `permits` over one second. * * This is equivalent to `AsyncMeter.newMeter(permits, 1.second, maxWaiters)`. + * + * @see [[perSecondUnbounded]] for creating an async meter that allows unbounded waiters */ def perSecond(permits: Int, maxWaiters: Int)(implicit timer: Timer): AsyncMeter = - newMeter(permits, 1.second, maxWaiters) + newMeter(permits, Duration.fromSeconds(1), maxWaiters) + + /** + * Creates an [[AsyncMeter]] that allows smoothed out `permits` per + * second, and has a maximum burst size of `permits` over one second. + * + * This meter is not bounded by the number of allowed waiters. This is equivalent to + * `AsyncMeter.newUnboundedMeter(permits, 1.second, maxWaiters)`. + * + * @see [[perSecond]] for creating an async meter that only allows a certain number of waiters + */ + def perSecondUnbounded(permits: Int)(implicit timer: Timer): AsyncMeter = + newUnboundedMeter(permits, Duration.fromSeconds(1)) /** * Creates an [[AsyncMeter]] that allows smoothed out `permits` per second, @@ -58,7 +71,7 @@ object AsyncMeter { * second, but the burst should be smoothed out after that. */ def perSecondLimited(permits: Int, maxWaiters: Int)(implicit timer: Timer): AsyncMeter = - newMeter(1, 1.second / permits, maxWaiters) + newMeter(1, Duration.fromSeconds(1) / permits, maxWaiters) /** * Creates an [[AsyncMeter]] that has a maximum burst size of `burstSize` over @@ -125,11 +138,7 @@ object AsyncMeter { val seqWithoutLast: Seq[Future[Unit]] = (0 until num).map(_ => meter.await(meter.burstSize)) val seq = if (last == 0) seqWithoutLast else seqWithoutLast :+ meter.await(last) val result = Future.join(seq) - result.onFailure { exc => - seq.foreach { f: Future[Unit] => - f.raise(exc) - } - } + result.onFailure { exc => seq.foreach { (f: Future[Unit]) => f.raise(exc) } } result } else meter.await(permits) } @@ -169,7 +178,8 @@ class AsyncMeter private ( private[concurrent] val burstSize: Int, burstDuration: Duration, q: BlockingQueue[(Promise[Unit], Int)] -)(implicit timer: Timer) { +)( + implicit timer: Timer) { require(burstSize > 0, s"burst size of $burstSize, which is <= 0 doesn't make sense") require( @@ -210,16 +220,16 @@ class AsyncMeter private ( * If the returned [[com.twitter.util.Future]] is interrupted, we * try to cancel it. If it's successfully cancelled, the * [[com.twitter.util.Future]] is satisfied with a - * [[java.util.concurrent.CancellationException]], and the permits + * `java.util.concurrent.CancellationException`, and the permits * will not be issued, so a subsequent waiter can take advantage * of the permits. * * If `await` is invoked when there are already `maxWaiters` waiters waiting * for permits, the [[com.twitter.util.Future]] is immediately satisfied with - * a [[java.util.concurrent.RejectedExecutionException]]. + * a `java.util.concurrent.RejectedExecutionException`. * * If more permits are requested than `burstSize` then it returns a failed - * [[java.lang.IllegalArgumentException]] [[com.twitter.util.Future]] + * `java.lang.IllegalArgumentException` [[com.twitter.util.Future]] * immediately. */ def await(permits: Int): Future[Unit] = { @@ -240,7 +250,7 @@ class AsyncMeter private ( // guarantees that satisfying the thread is not racy--we also use // Promise#setValue or Promise#setException to ensure that if there's a // race, it will fail loudly. - val p = Promise[Unit] + val p = Promise[Unit]() val tup = (p, permits) if (q.offer(tup)) { diff --git a/util-core/src/main/scala/com/twitter/concurrent/AsyncQueue.scala b/util-core/src/main/scala/com/twitter/concurrent/AsyncQueue.scala index b01156ca08..9333dc8174 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/AsyncQueue.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/AsyncQueue.scala @@ -1,15 +1,22 @@ package com.twitter.concurrent -import com.twitter.util.{Future, Promise, Try, Return, Throw} -import java.util.{Queue => JQueue, ArrayDeque} +import com.twitter.util.{Future, Promise, Return, Throw, Try} +import java.util.ArrayDeque import scala.collection.immutable.Queue object AsyncQueue { private sealed trait State - private case object Idle extends State - private case object Offering extends State + + // In the `OfferingOrIdle` state we can have 0 or more chunks of data ready for + // consumption. If we have 0 data chunks and the queue is `poll`ed we can transition + // to the `Polling` state. + private case object OfferingOrIdle extends State + + // If we're in a `Polling` state we must have at least one polling `Promise`. private case object Polling extends State + + // In the `Excepting` state the queue will only contain offers. private case class Excepting(exc: Throwable) extends State /** Indicates there is no max capacity */ @@ -18,8 +25,7 @@ object AsyncQueue { /** * An asynchronous FIFO queue. In addition to providing [[offer]] - * and [[poll]], the queue can be [[fail "failed"]], flushing current - * pollers. + * and [[poll]], the queue can be failed, flushing current pollers. * * @param maxPendingOffers optional limit on the number of pending `offers`. * The default is unbounded, but any other positive value can be used to limit @@ -33,17 +39,14 @@ class AsyncQueue[T](maxPendingOffers: Int) { require(maxPendingOffers > 0) - // synchronize all access to state, offers, and pollers - private[this] var state: State = Idle + // synchronize all access to `state` and `queue` + private[this] var state: State = OfferingOrIdle - // these aren't part of the state machine for performance - private[this] val offers: JQueue[T] = new ArrayDeque[T] - private[this] val pollers: JQueue[Promise[T]] = new ArrayDeque[Promise[T]] + private[this] val queue = new ArrayDeque[Any] /** * An asynchronous, unbounded, FIFO queue. In addition to providing [[offer]] - * and [[poll]], the queue can be [[fail "failed"]], flushing current - * pollers. + * and [[poll]], the queue can be failed, flushing current pollers. */ def this() = this(AsyncQueue.UnboundedCapacity) @@ -51,7 +54,10 @@ class AsyncQueue[T](maxPendingOffers: Int) { * Returns the current number of pending elements. */ final def size: Int = synchronized { - offers.size + state match { + case OfferingOrIdle | Excepting(_) => queue.size + case Polling => 0 // If we're `Polling` we're empty + } } /** @@ -60,28 +66,26 @@ class AsyncQueue[T](maxPendingOffers: Int) { */ final def poll(): Future[T] = synchronized { state match { - case Idle => + case OfferingOrIdle if queue.isEmpty => val p = new Promise[T] state = Polling - pollers.offer(p) + queue.offer(p) p + case OfferingOrIdle => + val elem = queue.poll() + Future.value(elem.asInstanceOf[T]) + case Polling => val p = new Promise[T] - pollers.offer(p) + queue.offer(p) p - case Offering => - val elem = offers.poll() - if (offers.isEmpty) - state = Idle - Future.value(elem) - - case Excepting(t) if offers.isEmpty => + case Excepting(t) if queue.isEmpty => Future.exception(t) case Excepting(_) => - Future.value(offers.poll()) + Future.value(queue.poll().asInstanceOf[T]) } } @@ -94,22 +98,17 @@ class AsyncQueue[T](maxPendingOffers: Int) { var waiter: Promise[T] = null val result = synchronized { state match { - case Idle => - state = Offering - offers.offer(elem) - true - - case Offering if offers.size >= maxPendingOffers => + case OfferingOrIdle if queue.size >= maxPendingOffers => false - case Offering => - offers.offer(elem) + case OfferingOrIdle => + queue.offer(elem) true case Polling => - waiter = pollers.poll() - if (pollers.isEmpty) - state = Idle + waiter = queue.poll().asInstanceOf[Promise[T]] + if (queue.isEmpty) + state = OfferingOrIdle true case Excepting(_) => @@ -126,28 +125,30 @@ class AsyncQueue[T](maxPendingOffers: Int) { /** * Drains any pending elements into a `Try[Queue]`. * - * If the queue has been [[fail failed]] and is now empty, + * If the queue has been failed and is now empty, * a `Throw` of the exception used to fail will be returned. * Otherwise, return a `Return(Queue)` of the pending elements. */ final def drain(): Try[Queue[T]] = synchronized { state match { - case Offering => - state = Idle + case OfferingOrIdle => var q = Queue.empty[T] - while (!offers.isEmpty) { - q :+= offers.poll() + while (!queue.isEmpty) { + q :+= queue.poll().asInstanceOf[T] } Return(q) - case Excepting(e) if !offers.isEmpty => + + case Excepting(_) if !queue.isEmpty => var q = Queue.empty[T] - while (!offers.isEmpty) { - q :+= offers.poll() + while (!queue.isEmpty) { + q :+= queue.poll().asInstanceOf[T] } Return(q) + case Excepting(e) => Throw(e) - case _ => + + case Polling => Return(Queue.empty) } } @@ -166,37 +167,39 @@ class AsyncQueue[T](maxPendingOffers: Int) { * No new elements are admitted to the queue after it has been failed. */ def fail(exc: Throwable, discard: Boolean): Unit = { - var q: Queue[Promise[T]] = null - synchronized { + // CAUTION: `q` may be `null` + val q: Array[AnyRef] = synchronized { state match { - case Idle => - state = Excepting(exc) - case Polling => state = Excepting(exc) - q = Queue.empty[Promise[T]] - while (!pollers.isEmpty) { - val waiter = pollers.poll() - q :+= waiter + if (queue.isEmpty) null + else { + val data = queue.toArray + queue.clear() + data } - case Offering => + case OfferingOrIdle => if (discard) - offers.clear() + queue.clear() state = Excepting(exc) + null case Excepting(_) => // Just take the first one. + null } } // we do this to avoid satisfaction while synchronized, which could lead to // lock contention if closures on the promise are slow or there are a lot of // them if (q != null) { - q.foreach { p => - p.setException(exc) + var i = 0 + while (i < q.length) { + q(i).asInstanceOf[Promise[_]].setException(exc) + i += 1 } } } - override def toString = "AsyncQueue<%s>".format(synchronized(state)) + override def toString: String = s"AsyncQueue<${synchronized(state)}>" } diff --git a/util-core/src/main/scala/com/twitter/concurrent/AsyncSemaphore.scala b/util-core/src/main/scala/com/twitter/concurrent/AsyncSemaphore.scala index 8e3ea27c62..51bce79d36 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/AsyncSemaphore.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/AsyncSemaphore.scala @@ -1,10 +1,9 @@ package com.twitter.concurrent import com.twitter.util._ -import java.util.ArrayDeque -import java.util.concurrent.RejectedExecutionException +import java.util.concurrent.{ConcurrentLinkedQueue, RejectedExecutionException} +import java.util.concurrent.atomic.AtomicInteger import scala.annotation.tailrec -import scala.collection.JavaConverters._ import scala.util.control.NonFatal /** @@ -17,7 +16,7 @@ import scala.util.control.NonFatal * {{{ * val semaphore = new AsyncSemaphore(n) * ... - * semaphore.acquireAndRun() { + * semaphore.acquireAndRun { * somethingThatReturnsFutureT() * } * }}} @@ -50,42 +49,48 @@ class AsyncSemaphore protected (initialPermits: Int, maxWaiters: Option[Int]) { require(maxWaiters.getOrElse(0) >= 0, s"maxWaiters must be non-negative: $maxWaiters") require(initialPermits > 0, s"initialPermits must be positive: $initialPermits") - // access to `closed`, `waitq`, and `availablePermits` is synchronized - // by locking on `lock` - private[this] var closed: Option[Throwable] = None - private[this] val waitq = new ArrayDeque[Promise[Permit]] - private[this] var availablePermits = initialPermits + // if failedState (Int.MinValue), failed + // else if >= 0, # of available permits + // else if < 0, # of waiters + private[this] final val state = new AtomicInteger(initialPermits) + @volatile private[this] final var failure: Future[Nothing] = null - // Serves as our intrinsic lock. - private[this] final def lock: Object = waitq + private[this] final val failedState = Int.MinValue + private[this] final val waiters = new ConcurrentLinkedQueue[Waiter] + private[this] final val waitersLimit = maxWaiters.map(v => -v).getOrElse(failedState + 1) - private[this] val semaphorePermit = new Permit { - private[this] val ReturnThis = Return(this) - - @tailrec override def release(): Unit = { - val waiter = lock.synchronized { - val next = waitq.pollFirst() - if (next == null) { - availablePermits += 1 - } - next - } - - if (waiter != null) { + private[this] final val permit = new Permit { + override final def release(): Unit = AsyncSemaphore.this.release() + } + private[this] final val permitFuture = Future.value(permit) + private[this] final val permitReturn = Return(permit) + private[this] final val releasePermit: Any => Unit = _ => permit.release() + + @tailrec private[this] final def release(): Unit = { + val s = state.get() + if (s != failedState) { + if (!state.compareAndSet(s, s + 1) || // If the waiter is already satisfied, it must // have been interrupted, so we can simply move on. - if (!waiter.updateIfEmpty(ReturnThis)) { - release() - } + (s < 0 && !pollWaiter().updateIfEmpty(permitReturn))) { + release() } } } - private[this] val futurePermit = Future.value(semaphorePermit) + def numWaiters: Int = { + val s = state.get() + if (s < 0 && s != failedState) -s + else 0 + } - def numWaiters: Int = lock.synchronized(waitq.size) def numInitialPermits: Int = initialPermits - def numPermitsAvailable: Int = lock.synchronized(availablePermits) + + def numPermitsAvailable: Int = { + val s = state.get() + if (s > 0) s + else 0 + } /** * Fail the semaphore and stop it from distributing further permits. Subsequent @@ -93,12 +98,29 @@ class AsyncSemaphore protected (initialPermits: Int, maxWaiters: Option[Int]) { * are also failed with `exc`. */ def fail(exc: Throwable): Unit = { - val drained = lock.synchronized { - closed = Some(exc) - waitq.asScala.toList + failure = Future.exception(exc) + var s = state.getAndSet(failedState) + // drain the wait queue + while (s < 0 && s != failedState) { + pollWaiter().raise(exc) + s += 1 + } + } + + /** + * Polls from the wait queue until it returns an item. Must be called + * after a state modification that indicates that a new waiter will be available + * shortly. The retry accounts for the race when `acquire` updates the state to + * indicate that a waiter is available but it hasn't been added to the queue yet. + * @return + */ + @tailrec private[this] final def pollWaiter(): Waiter = { + val w = waiters.poll() + if (w != null) { + w + } else { + pollWaiter() } - // delegate dequeuing to the interrupt handler defined in #acquire - drained.foreach(_.raise(exc)) } /** @@ -116,29 +138,26 @@ class AsyncSemaphore protected (initialPermits: Int, maxWaiters: Option[Int]) { * or a Future.Exception[RejectedExecutionException]` if the configured maximum * number of waiters would be exceeded. */ - def acquire(): Future[Permit] = lock.synchronized { - if (closed.isDefined) - return Future.exception(closed.get) - - if (availablePermits > 0) { - availablePermits -= 1 - futurePermit - } else { - maxWaiters match { - case Some(max) if waitq.size >= max => - MaxWaitersExceededException - case _ => - val promise = new Promise[Permit] - promise.setInterruptHandler { - case t: Throwable => - if (promise.updateIfEmpty(Throw(t))) lock.synchronized { - waitq.remove(promise) - } - } - waitq.addLast(promise) - promise + def acquire(): Future[Permit] = { + @tailrec def loop(): Future[Permit] = { + val s = state.get() + if (s == failedState) { + failure + } else if (s == waitersLimit) { + AsyncSemaphore.MaxWaitersExceededException + } else if (state.compareAndSet(s, s - 1)) { + if (s > 0) { + permitFuture + } else { + val w = new Waiter + waiters.add(w) + w + } + } else { + loop() } } + loop() } /** @@ -153,17 +172,16 @@ class AsyncSemaphore protected (initialPermits: Int, maxWaiters: Option[Int]) { */ def acquireAndRun[T](func: => Future[T]): Future[T] = acquire().flatMap { permit => - val f = try func - catch { - case NonFatal(e) => - Future.exception(e) - case e: Throwable => - permit.release() - throw e - } - f.ensure { - permit.release() - } + val f = + try func + catch { + case NonFatal(e) => + Future.exception(e) + case e: Throwable => + permit.release() + throw e + } + f.respond(releasePermit) } /** @@ -177,14 +195,22 @@ class AsyncSemaphore protected (initialPermits: Int, maxWaiters: Option[Int]) { * returned. */ def acquireAndRunSync[T](func: => T): Future[T] = - acquire().flatMap { permit => - Future(func).ensure { + acquire().map { _ => + try { + func + } finally { permit.release() } } } object AsyncSemaphore { + + private final class Waiter extends Promise[Permit] with Promise.InterruptHandler { + // the waiter will be removed from the queue later when a permit is released + override protected def onInterrupt(t: Throwable): Unit = updateIfEmpty(Throw(t)) + } + private val MaxWaitersExceededException = Future.exception(new RejectedExecutionException("Max waiters exceeded")) } diff --git a/util-core/src/main/scala/com/twitter/concurrent/AsyncStream.scala b/util-core/src/main/scala/com/twitter/concurrent/AsyncStream.scala index 8858313d96..a26def1ace 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/AsyncStream.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/AsyncStream.scala @@ -1,9 +1,8 @@ package com.twitter.concurrent -import com.twitter.io.{Buf, Reader} -import com.twitter.util.{Future, Promise, Return, Throw, Try} +import com.twitter.conversions.SeqUtil +import com.twitter.util.{Future, Return, Throw, Promise} import scala.annotation.varargs -import scala.collection.mutable /** * A representation of a lazy (and possibly infinite) sequence of asynchronous @@ -56,7 +55,7 @@ sealed abstract class AsyncStream[+A] { */ def tail: Future[Option[AsyncStream[A]]] = this match { case Empty | FromFuture(_) => Future.None - case Cons(_, more) => Future.value(Some(more())) + case Cons(_, more) => extract(more()) case Embed(fas) => fas.flatMap(_.tail) } @@ -80,9 +79,7 @@ sealed abstract class AsyncStream[+A] { * Note: forces the stream. For infinite streams, the future never resolves. */ def foreach(f: A => Unit): Future[Unit] = - foldLeft(()) { (_, a) => - f(a) - } + foldLeft(()) { (_, a) => f(a) } /** * Execute the specified effect as each element of the resulting @@ -98,9 +95,7 @@ sealed abstract class AsyncStream[+A] { * whether the entire stream is consumed. */ def withEffect(f: A => Unit): AsyncStream[A] = - map { a => - f(a); a - } + map { a => f(a); a } /** * Maps each element of the stream to a Future action, resolving them from @@ -110,22 +105,22 @@ sealed abstract class AsyncStream[+A] { * Note: forces the stream. For infinite streams, the future never resolves. */ def foreachF(f: A => Future[Unit]): Future[Unit] = - foldLeftF(()) { (_, a) => - f(a) - } + foldLeftF(()) { (_, a) => f(a) } /** - * Map over this stream with the given concurrency. The items will - * likely not be processed in order. `concurrencyLevel` specifies an - * "eagerness factor", and that many actions will be started when this - * method is called. Forcing the stream will yield the results of - * completed actions, and will block if none of the actions has yet - * completed. + * Map over this stream with the given concurrency. The items will likely be + * completed out of order, and if so, the stream of results will also be out + * of order. * - * This method is useful for speeding up calculations over a stream - * where executing the actions in order is not important. To implement - * a concurrent fold, first call `mapConcurrent` and then fold that - * stream. Similarly, concurrent `foreachF` can be achieved by + * `concurrencyLevel` specifies an "eagerness factor", and that many actions + * will be started when this method is called. Forcing the stream will yield + * the results of completed actions, returning the results that finished + * first, and will block if none of the actions has yet completed. + * + * This method is useful for speeding up calculations over a stream where + * executing the actions (and processing their results) in order is not + * important. To implement a concurrent fold, first call `mapConcurrent` and + * then fold that stream. Similarly, concurrent `foreachF` can be achieved by * applying `mapConcurrent` and then `foreach`. * * @param concurrencyLevel: How many actions to execute concurrently. This @@ -133,15 +128,8 @@ sealed abstract class AsyncStream[+A] { * the next item available in the stream. */ def mapConcurrent[B](concurrencyLevel: Int)(f: A => Future[B]): AsyncStream[B] = { - if (concurrencyLevel == 1) { - mapF(f) - } else if (concurrencyLevel < 1) { - throw new IllegalArgumentException( - s"concurrencyLevel must be at least one. got: $concurrencyLevel" - ) - } else { - embed(AsyncStream.mapConcStep(concurrencyLevel, f, Nil, () => this)) - } + require(concurrencyLevel > 0, s"concurrencyLevel must be at least one. got: $concurrencyLevel") + merge(fanout(this, concurrencyLevel).map(_.mapF(f)): _*) } /** @@ -158,9 +146,7 @@ sealed abstract class AsyncStream[+A] { this match { case Empty => empty case FromFuture(fa) => - Embed(fa.map { a => - if (p(a)) this else empty - }) + Embed(fa.map { a => if (p(a)) this else empty }) case Cons(fa, more) => Embed(fa.map { a => if (p(a)) Cons(fa, () => more().takeWhile(p)) @@ -182,9 +168,7 @@ sealed abstract class AsyncStream[+A] { this match { case Empty => empty case FromFuture(fa) => - Embed(fa.map { a => - if (p(a)) empty else this - }) + Embed(fa.map { a => if (p(a)) empty else this }) case Cons(fa, more) => Embed(fa.map { a => if (p(a)) more().dropWhile(p) @@ -223,13 +207,18 @@ sealed abstract class AsyncStream[+A] { /** * Map a function `f` over the elements in this stream and concatenate the * results. + * + * @note We use `flatMap` on `Future` instead of `map` to maintain + * stack-safety for monadic recursion. */ def flatMap[B](f: A => AsyncStream[B]): AsyncStream[B] = this match { case Empty => empty - case FromFuture(fa) => Embed(fa.map(f)) - case Cons(fa, more) => Embed(fa.map(f)) ++ more().flatMap(f) - case Embed(fas) => Embed(fas.map(_.flatMap(f))) + case FromFuture(fa) => Embed(fa.flatMap(a => Future.value(f(a)))) + case Cons(fa, more) => + Embed(fa.flatMap(a => Future.value(f(a)))) ++ more().flatMap(f) + case Embed(fas) => + Embed(fas.flatMap(as => Future.value(as.flatMap(f)))) } /** @@ -255,9 +244,7 @@ sealed abstract class AsyncStream[+A] { this match { case Empty => empty case FromFuture(fa) => - Embed(fa.map { a => - if (p(a)) this else empty - }) + Embed(fa.map { a => if (p(a)) this else empty }) case Cons(fa, more) => Embed(fa.map { a => if (p(a)) Cons(fa, () => more().filter(p)) @@ -350,12 +337,12 @@ sealed abstract class AsyncStream[+A] { } /** - * Helper method used to avoid scanLeft being one behind in case of Embed AsyncStream. - * - * scanLeftEmbed, unlike scanLeft, does not return the initial value `z` and is there to - * prevent the Embed case from returning duplicate initial `z` values for scanLeft. - * - */ + * Helper method used to avoid scanLeft being one behind in case of Embed AsyncStream. + * + * scanLeftEmbed, unlike scanLeft, does not return the initial value `z` and is there to + * prevent the Embed case from returning duplicate initial `z` values for scanLeft. + * + */ private def scanLeftEmbed[B](z: B)(f: (B, A) => B): AsyncStream[B] = this match { case Embed(fas) => Embed(fas.map(_.scanLeftEmbed(z)(f))) @@ -469,7 +456,7 @@ sealed abstract class AsyncStream[+A] { * Note: forces the stream. For infinite streams, the future never resolves. */ def observe(): Future[(Seq[A], Option[Throwable])] = { - val buf: mutable.ListBuffer[A] = mutable.ListBuffer.empty + val buf = Vector.newBuilder[A] def go(as: AsyncStream[A]): Future[Unit] = as match { @@ -489,8 +476,8 @@ sealed abstract class AsyncStream[+A] { } go(this).transform { - case Throw(exc) => Future.value(buf.toList -> Some(exc)) - case Return(_) => Future.value(buf.toList -> None) + case Throw(exc) => Future.value(buf.result() -> Some(exc)) + case Return(_) => Future.value(buf.result() -> None) } } @@ -498,25 +485,28 @@ sealed abstract class AsyncStream[+A] { * Buffer the specified number of items from the stream, or all * remaining items if the end of the stream is reached before finding * that many items. In all cases, this method should act like - * + * * and not cause evaluation of the remainder of the stream. */ private[concurrent] def buffer(n: Int): Future[(Seq[A], () => AsyncStream[A])] = { // pre-allocate the buffer, unless it's very large - val buffer = new mutable.ArrayBuffer[A](n.min(1024)) + val buffer = Vector.newBuilder[A] + buffer.sizeHint(n.max(0).min(1024)) def fillBuffer( sizeRemaining: Int - )(s: => AsyncStream[A]): Future[(Seq[A], () => AsyncStream[A])] = - if (sizeRemaining < 1) Future.value((buffer, () => s)) + )( + s: => AsyncStream[A] + ): Future[(Seq[A], () => AsyncStream[A])] = + if (sizeRemaining < 1) Future.value((buffer.result(), () => s)) else s match { - case Empty => Future.value((buffer, () => s)) + case Empty => Future.value((buffer.result(), () => s)) case FromFuture(fa) => fa.flatMap { a => buffer += a - Future.value((buffer, () => empty)) + Future.value((buffer.result(), () => empty)) } case Cons(fa, more) => @@ -536,7 +526,7 @@ sealed abstract class AsyncStream[+A] { * Convert the stream into a stream of groups of items. This * facilitates batch processing of the items in the stream. In all * cases, this method should act like - * + * * The resulting stream will cause this original stream to be * evaluated group-wise, so calling this method will cause the first * `groupSize` cells to be evaluated (even without examining the @@ -587,27 +577,28 @@ sealed abstract class AsyncStream[+A] { * stream to occur, but do not need to do anything with the resulting * values. */ - def force: Future[Unit] = foreach { _ => - } + def force: Future[Unit] = foreach { _ => } } object AsyncStream { private case object Empty extends AsyncStream[Nothing] private case class Embed[A](fas: Future[AsyncStream[A]]) extends AsyncStream[A] private case class FromFuture[A](fa: Future[A]) extends AsyncStream[A] - private class Cons[A](val fa: Future[A], next: () => AsyncStream[A]) extends AsyncStream[A] { - private[this] lazy val _more: AsyncStream[A] = next() - def more(): AsyncStream[A] = _more - override def toString(): String = s"Cons($fa, $next)" - } + /** + * This is defined as a case class so that we get a generated unapply that works with + * [[AsyncStream]]s covariant type in Scala 3. See CSL-11207 or + * https://github.com/lampepfl/dotty/issues/13125 + */ + private case class Cons[A] private (fa: Future[A], next: () => AsyncStream[A]) + extends AsyncStream[A] object Cons { - def apply[A](fa: Future[A], next: () => AsyncStream[A]): AsyncStream[A] = - new Cons(fa, next) - - def unapply[A](as: Cons[A]): Option[(Future[A], () => AsyncStream[A])] = - // note: pattern match returns the memoized value - Some((as.fa, () => as.more())) + def apply[A](fut: Future[A], nextAsync: () => AsyncStream[A]): AsyncStream[A] = { + lazy val _more: AsyncStream[A] = nextAsync() + new Cons(fut, () => _more) { + override def toString(): String = s"Cons($fut, $nextAsync)" + } + } } implicit class Ops[A](tail: => AsyncStream[A]) { @@ -629,11 +620,8 @@ object AsyncStream { * {{{ * AsyncStream(1,2,3) * }}} - * - * Note: we can't annotate this with varargs because of - * https://issues.scala-lang.org/browse/SI-8743. This seems to be fixed in a - * more recent scala version so we will revisit this soon. */ + @varargs def apply[A](as: A*): AsyncStream[A] = fromSeq(as) /** @@ -654,22 +642,22 @@ object AsyncStream { fromFuture[A](Future.exception(e)) /** - * Transformation (or lift) from [[Seq]] into `AsyncStream`. + * Transformation (or lift) from `Seq` into `AsyncStream`. */ def fromSeq[A](seq: Seq[A]): AsyncStream[A] = seq match { case Nil => empty - case _ if seq.hasDefiniteSize && seq.tail.isEmpty => of(seq.head) + case _ if SeqUtil.hasKnownSize(seq) && seq.tail.isEmpty => of(seq.head) case _ => seq.head +:: fromSeq(seq.tail) } /** - * Transformation (or lift) from [[Future]] into `AsyncStream`. + * Transformation (or lift) from [[com.twitter.util.Future]] into `AsyncStream`. */ def fromFuture[A](f: Future[A]): AsyncStream[A] = FromFuture(f) /** - * Transformation (or lift) from [[Option]] into `AsyncStream`. + * Transformation (or lift) from `Option` into `AsyncStream`. */ def fromOption[A](o: Option[A]): AsyncStream[A] = o match { @@ -677,21 +665,19 @@ object AsyncStream { case Some(a) => of(a) } - /** - * Transformation (or lift) from [[Reader]] into `AsyncStream`. - */ - def fromReader[A <: Buf](r: Reader[A], chunkSize: Int = Int.MaxValue): AsyncStream[A] = - fromFuture(r.read(chunkSize)).flatMap { - case Some(buf) => buf +:: fromReader(r, chunkSize) - case None => AsyncStream.empty[A] - } - /** * Lift from [[Future]] into `AsyncStream` and then flatten. */ private[concurrent] def embed[A](fas: Future[AsyncStream[A]]): AsyncStream[A] = Embed(fas) + private def extract[A](as: AsyncStream[A]): Future[Option[AsyncStream[A]]] = + as match { + case Empty => Future.None + case Embed(fas) => fas.flatMap(extract) + case _ => Future.value(Some(as)) + } + /** * Java friendly [[AsyncStream.flatten]]. */ @@ -727,65 +713,108 @@ object AsyncStream { } } - /** - * This must exist here, otherwise scalac will be greedy and hold a reference to the head of the stream - */ - private def mapConcStep[A, B]( - concurrencyLevel: Int, - f: (A => Future[B]), - pending: Seq[Future[Option[B]]], - inputs: () => AsyncStream[A] - ): Future[AsyncStream[B]] = { - // We only invoke inputs().uncons if there is space for more work - // to be started. - val inputReady: Future[Left[Option[(A, () => AsyncStream[A])], Nothing]] = - if (pending.size >= concurrencyLevel) Future.never - else - inputs().uncons.flatMap { - case None if pending.nonEmpty => Future.never - case other => Future.value(Left(other)) + // Oneshot turns an AsyncStream into a resource, where each read depletes the + // stream of that item. In other words, each item of the stream is observable + // only once. + // + // more must be guarded by synchronization on this. + private final class Oneshot[A](var more: () => AsyncStream[A]) { + def read(): Future[Option[A]] = { + // References to readEmbed args for the Embed case. + var fas: Future[AsyncStream[A]] = null + var headp: Promise[Option[A]] = null + var tailp: Promise[() => AsyncStream[A]] = null + + // Reference to head for the default case. + var headf: Future[Option[A]] = null + + synchronized { + more() match { + case Embed(embedded) => + // The Embed case is special because + // + // val head = as.head + // as = as.drop(1) + // + // reduces to a chain of maps, which can be very long, as many as the + // number of times read is called while waiting for the embedded + // AsyncStream. + // + // val head = fas.map(_.drop(1)).map(_.drop(1))....flatMap(_.head) + // as = Embed(fas.map(_.drop(1)).map(_.drop(1))...) + // + // In other words, for AsyncStreams with Embed tails, the naive + // implementation multiplies the work of each item by the concurrency + // level. + // + // The optimization here is straight forward: create a head and tail + // promise and that wait for the embedded AsyncStream. Subsequent + // calls to read will create new head and tail promises that wait on + // the previous tail promise. + fas = embedded + headp = new Promise[Option[A]] + tailp = new Promise[() => AsyncStream[A]] + more = () => AsyncStream.embed(tailp.map(_())) + headp + + case as => + headf = as.head + more = () => as.drop(1) } + } - // The inputReady.isDefined check is an optimization to avoid the - // wasted allocation of calling Future.select when we know that - // there is more work that is ready to be started. - val workDone: Future[Right[Nothing, (Try[Option[B]], Seq[Future[Option[B]]])]] = - if (pending.isEmpty || inputReady.isDefined) Future.never - else Future.select(pending).map(Right(_)) - - // Wait for either the next input to be ready (Left) or for some pending - // work to complete (Right). - inputReady.or(workDone).flatMap { - case Left(None) => - // There is no work pending and we have exhausted the - // inputs, so we are done. - Future.value(empty) - - case Left(Some((a, tl))) => - // There is available concurrency and a new input is ready, - // so start the work and add it to `pending`. - mapConcStep(concurrencyLevel, f, f(a).map(Some(_)) +: pending, tl) - - case Right((Throw(t), _)) => - // Some work finished with failure, so terminate the stream. - Future.exception(t) - - case Right((Return(None), newPending)) => - // A cons cell was forced, freeing up a spot in `pending`. - mapConcStep(concurrencyLevel, f, newPending, inputs) - - case Right((Return(Some(a)), newPending)) => - // Some work finished. Eagerly start the next step, so - // that all available inputs have work started on them. In - // the next step, we replace the pending Future for the - // work that just completed with a Future that will be - // satisfied by forcing the next element of the result - // stream. Keeping `pending` full is the mechanism that we - // use to bound the evaluation depth at `concurrencyLevel` - // until the stream is forced. - val cellForced = new Promise[Option[B]] - val rest = mapConcStep(concurrencyLevel, f, cellForced +: newPending, inputs) - Future.value(mk(a, { cellForced.setValue(None); embed(rest) })) + // Non-null fas implies Embed case. + if (fas != null) { + readEmbed(fas, headp, tailp) + headp + } else { + headf + } } + + private[this] def readEmbed( + fas: Future[AsyncStream[A]], + headp: Promise[Option[A]], + tailp: Promise[() => AsyncStream[A]] + ): Unit = + fas.respond { + case Throw(e) => + headp.setException(e) + tailp.setException(e) + case Return(v) => + v match { + case Empty => + headp.become(Future.None) + tailp.setValue(Oneshot.empty) + case FromFuture(fa) => + headp.become(fa.map(Some(_))) + tailp.setValue(Oneshot.empty) + case Cons(fa, more2) => + headp.become(fa.map(Some(_))) + tailp.setValue(more2) + case Embed(newFas) => + readEmbed(newFas, headp, tailp) + } + } + + // Recover the AsyncStream interface. We don't return `as` because we want + // to deplete the original stream with each read. With this we can create + // multiple AsyncStream views of the same Oneshot resource, thus "fanning + // out" the original stream into multiple distinct AsyncStreams. + def toAsyncStream: AsyncStream[A] = + AsyncStream.fromFuture(read()).flatMap { + case None => AsyncStream.empty + case Some(a) => a +:: toAsyncStream + } + } + + private object Oneshot { + val emptyVal: () => AsyncStream[Nothing] = () => AsyncStream.empty[Nothing] + def empty[A]: () => AsyncStream[A] = emptyVal.asInstanceOf[() => AsyncStream[A]] + } + + private def fanout[A](as: AsyncStream[A], n: Int): Seq[AsyncStream[A]] = { + val reader = new Oneshot(() => as) + 0.until(n).map(_ => reader.toAsyncStream) } } diff --git a/util-core/src/main/scala/com/twitter/concurrent/BUILD b/util-core/src/main/scala/com/twitter/concurrent/BUILD new file mode 100644 index 0000000000..0c1d94d45f --- /dev/null +++ b/util-core/src/main/scala/com/twitter/concurrent/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-core-concurrent", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-core/src/main/scala/com/twitter/concurrent/Broker.scala b/util-core/src/main/scala/com/twitter/concurrent/Broker.scala index 937c416c41..991bd422d9 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Broker.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/Broker.scala @@ -73,49 +73,56 @@ class Broker[T] { Future.value(sendTx) } - case s @ (Quiet | Sending(_)) => - val p = new Promise[Tx[Unit]] - val elem: (Promise[Tx[Unit]], T) = (p, msg) - p.setInterruptHandler { - case _ => rmElem(elem) - } - val nextState = s match { - case Quiet => Sending(Queue(elem)) - case Sending(q) => Sending(q enqueue elem) - case Receiving(_) => throw new IllegalStateException() - } - - if (state.compareAndSet(s, nextState)) p else prepare() + case s @ Sending(q) => + val elem = createElem() + if (state.compareAndSet(s, Sending(q enqueue elem))) elem._1 + else prepare() + + case Quiet => + val elem = createElem() + if (state.compareAndSet(Quiet, Sending(Queue(elem)))) elem._1 + else prepare() } } + + def createElem(): (Promise[Tx[Unit]], T) = { + val p = new Promise[Tx[Unit]] + val elem = (p, msg) + p.setInterruptHandler { case _ => rmElem(elem) } + elem + } } val recv: Offer[T] = new Offer[T] { @tailrec - def prepare() = - state.get match { - case s @ Sending(sq) => - if (sq.isEmpty) throw new IllegalStateException() - val ((sendp, msg), newq) = sq.dequeue - val nextState = if (newq.isEmpty) Quiet else Sending(newq) - if (!state.compareAndSet(s, nextState)) prepare() - else { - val (sendTx, recvTx) = Tx.twoParty(msg) - sendp.setValue(sendTx) - Future.value(recvTx) - } + def prepare() = state.get match { + case s @ Sending(sq) => + if (sq.isEmpty) throw new IllegalStateException() + val ((sendp, msg), newq) = sq.dequeue + val nextState = if (newq.isEmpty) Quiet else Sending(newq) + if (!state.compareAndSet(s, nextState)) prepare() + else { + val (sendTx, recvTx) = Tx.twoParty(msg) + sendp.setValue(sendTx) + Future.value(recvTx) + } - case s @ (Quiet | Receiving(_)) => - val p = new Promise[Tx[T]] - p.setInterruptHandler { case _ => rmElem(p) } - val nextState = s match { - case Quiet => Receiving(Queue(p)) - case Receiving(q) => Receiving(q enqueue p) - case Sending(_) => throw new IllegalStateException() - } + case s @ Receiving(q) => + val p = createPromise() + if (state.compareAndSet(s, Receiving(q enqueue p))) p + else prepare() + + case Quiet => + val p = createPromise() + if (state.compareAndSet(Quiet, Receiving(Queue(p)))) p + else prepare() + } - if (state.compareAndSet(s, nextState)) p else prepare() - } + def createPromise(): Promise[Tx[T]] = { + val p = new Promise[Tx[T]] + p.setInterruptHandler { case _ => rmElem(p) } + p + } } /* Scala actor style / CSP syntax. */ diff --git a/util-core/src/main/scala/com/twitter/concurrent/Serialized.scala b/util-core/src/main/scala/com/twitter/concurrent/Serialized.scala index a9f6ac22a4..1dd27ec135 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Serialized.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/Serialized.scala @@ -19,7 +19,7 @@ trait Serialized { } } - private[this] val nwaiters = new AtomicInteger(0) + private[this] val nwaiters: AtomicInteger = new AtomicInteger(0) protected val serializedQueue: java.util.Queue[Job[_]] = new ConcurrentLinkedQueue[Job[_]] protected def serialized[A](f: => A): Future[A] = { @@ -28,9 +28,10 @@ trait Serialized { serializedQueue add { Job(result, () => f) } if (nwaiters.getAndIncrement() == 0) { - do { + Try { serializedQueue.remove()() } + while (nwaiters.decrementAndGet() > 0) { Try { serializedQueue.remove()() } - } while (nwaiters.decrementAndGet() > 0) + } } result diff --git a/util-core/src/main/scala/com/twitter/concurrent/Spool.scala b/util-core/src/main/scala/com/twitter/concurrent/Spool.scala index 2ed1d09c7d..90caec8c0c 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/Spool.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/Spool.scala @@ -1,14 +1,14 @@ package com.twitter.concurrent import com.twitter.util.{Await, Duration, Future, Return, Throw, ConstFuture} -import scala.collection.mutable.ArrayBuffer import scala.collection.mutable import scala.language.implicitConversions import java.io.EOFException /** * Note: [[Spool]] is no longer the recommended asynchronous stream abstraction. - * We encourage you to use [[AsyncStream]] instead. + * We encourage you to use [[AsyncStream]] instead. See [[Spools]] for the Java + * adaptation of the [[Spool]] companion object. * * A spool is an asynchronous stream. It more or less mimics the scala * {{Stream}} collection, but with cons cells that have either eager or @@ -36,8 +36,6 @@ import java.io.EOFException * fill(rest) * firstElem *:: rest * }}} - * - * Note: There is a Java-friendly API for this trait: [[com.twitter.concurrent.AbstractSpool]]. */ sealed trait Spool[+A] { // NB: Spools are always lazy internally in order to provide the expected behavior @@ -67,7 +65,7 @@ sealed trait Spool[+A] { * Apply {{f}} for each item in the spool, until the end. {{f}} is * applied as the items become available. */ - def foreach[B](f: A => B) = foreachElem(_ foreach f) + def foreach[B](f: A => B): Future[Unit] = foreachElem(_ foreach f) /** * A version of {{foreach}} that wraps each element in an @@ -79,12 +77,12 @@ sealed trait Spool[+A] { Future { f(Some(head)) } flatMap { _ => tail transform { case Return(s) => s.foreachElem(f) - case Throw(_: EOFException) => Future { f(None) } + case Throw(_: EOFException) => Future { f(None); () } case Throw(cause) => Future.exception(cause) } } } else { - Future { f(None) } + Future { f(None); () } } } @@ -161,9 +159,7 @@ sealed trait Spool[+A] { def mapFuture[B](f: A => Future[B]): Future[Spool[B]] = { if (isEmpty) Future.value(empty[B]) else { - f(head) map { h => - new LazyCons(h, tail flatMap (_ mapFuture f)) - } + f(head) map { h => new LazyCons(h, tail flatMap (_ mapFuture f)) } } } @@ -261,28 +257,19 @@ sealed trait Spool[+A] { * satisfied when the entire result is ready. */ def toSeq: Future[Seq[A]] = { - val as = new ArrayBuffer[A] + val as = Vector.newBuilder[A] foreach { a => as += a - } map { _ => - as - } + } map { _ => as.result() } } /** * Eagerly executes all computation represented by this Spool (presumably for * side-effects), and returns a Future representing its completion. */ - def force: Future[Unit] = foreach { _ => - () - } + def force: Future[Unit] = foreach { _ => () } } -/** - * Abstract `Spool` class for Java compatibility. - */ -abstract class AbstractSpool[A] extends Spool[A] - /** * Note: [[Spool]] is no longer the recommended asynchronous stream abstraction. * We encourage you to use [[AsyncStream]] instead. @@ -291,22 +278,22 @@ abstract class AbstractSpool[A] extends Spool[A] */ object Spool { case class Cons[A](head: A, tail: Future[Spool[A]]) extends Spool[A] { - def isEmpty = false - override def toString = "Cons(%s, %c)".format(head, if (tail.isDefined) '*' else '?') + def isEmpty: Boolean = false + override def toString: String = "Cons(%s, %c)".format(head, if (tail.isDefined) '*' else '?') } private class LazyCons[A](val head: A, next: => Future[Spool[A]]) extends Spool[A] { def isEmpty = false - lazy val tail = next + lazy val tail: Future[Spool[A]] = next // NB: not touching tail, to avoid forcing unnecessarily - override def toString = "Cons(%s, ?)".format(head) + override def toString: String = "Cons(%s, ?)".format(head) } object Empty extends Spool[Nothing] { - def isEmpty = true - def head = throw new NoSuchElementException("spool is empty") - def tail = Future.exception(new NoSuchElementException("spool is empty")) - override def toString = "Empty" + def isEmpty: Boolean = true + def head: Nothing = throw new NoSuchElementException("spool is empty") + def tail: Future[Nothing] = Future.exception(new NoSuchElementException("spool is empty")) + override def toString: String = "Empty" } /** @@ -352,7 +339,7 @@ object Spool { * @deprecated Deprecated in favor of {{*::}}. This will eventually be removed. */ @deprecated("Use *:: instead.", "6.14.1") - def **::(head: A) = cons(head, tail) + def **::(head: A): Spool[A] = cons(head, tail) } implicit def syntax1[A](s: Spool[A]): Syntax1[A] = new Syntax1[A](s) @@ -384,7 +371,7 @@ object Spool { */ class ToSpool[A](s: Seq[A]) { def toSpool: Spool[A] = - s.reverse.foldLeft(Spool.empty: Spool[A]) { + s.reverseIterator.foldLeft(Spool.empty: Spool[A]) { case (tail, head) => head *:: Future.value(tail) } } diff --git a/util-core/src/main/scala/com/twitter/concurrent/SpoolSource.scala b/util-core/src/main/scala/com/twitter/concurrent/SpoolSource.scala index f7892edc80..e11160a4fc 100644 --- a/util-core/src/main/scala/com/twitter/concurrent/SpoolSource.scala +++ b/util-core/src/main/scala/com/twitter/concurrent/SpoolSource.scala @@ -9,7 +9,7 @@ import com.twitter.util.{Future, Promise, Return} object SpoolSource { private object DefaultInterruptHandler extends PartialFunction[Any, Nothing] { def isDefinedAt(x: Any) = false - def apply(x: Any) = throw new MatchError(x) + def apply(x: Any): Nothing = throw new MatchError(x) } } @@ -106,7 +106,11 @@ class SpoolSource[A](interruptHandler: PartialFunction[Throwable, Unit]) { } @tailrec - private[this] def updatingTailCall(newPromise: Promise[Spool[A]])(f: Promise[Spool[A]] => Unit): Unit = { + private[this] def updatingTailCall( + newPromise: Promise[Spool[A]] + )( + f: Promise[Spool[A]] => Unit + ): Unit = { val currentPromise = promiseRef.get // if the current promise is emptyPromise, then this source has already been closed if (currentPromise ne emptyPromise) { diff --git a/util-core/src/main/scala/com/twitter/conversions/BUILD b/util-core/src/main/scala/com/twitter/conversions/BUILD new file mode 100644 index 0000000000..d45b717d2e --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/BUILD @@ -0,0 +1,17 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-core-conversions", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], + exports = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-core/src/main/scala/com/twitter/conversions/DurationOps.scala b/util-core/src/main/scala/com/twitter/conversions/DurationOps.scala new file mode 100644 index 0000000000..ea104ce058 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/DurationOps.scala @@ -0,0 +1,52 @@ +package com.twitter.conversions + +import com.twitter.util.Duration +import java.util.concurrent.TimeUnit +import scala.language.implicitConversions + +/** + * Implicits for writing readable [[com.twitter.util.Duration]]s. + * + * @example + * {{{ + * import com.twitter.conversions.DurationOps._ + * + * 2000.nanoseconds + * 50.milliseconds + * 1.second + * 24.hours + * 40.days + * }}} + */ +object DurationOps { + + implicit class RichDuration(val numNanos: Long) extends AnyVal { + def nanoseconds: Duration = Duration(numNanos, TimeUnit.NANOSECONDS) + def nanosecond: Duration = nanoseconds + def microseconds: Duration = Duration(numNanos, TimeUnit.MICROSECONDS) + def microsecond: Duration = microseconds + def milliseconds: Duration = Duration(numNanos, TimeUnit.MILLISECONDS) + def millisecond: Duration = milliseconds + def millis: Duration = milliseconds + def seconds: Duration = Duration(numNanos, TimeUnit.SECONDS) + def second: Duration = seconds + def minutes: Duration = Duration(numNanos, TimeUnit.MINUTES) + def minute: Duration = minutes + def hours: Duration = Duration(numNanos, TimeUnit.HOURS) + def hour: Duration = hours + def days: Duration = Duration(numNanos, TimeUnit.DAYS) + def day: Duration = days + } + + /** + * Forwarder for Int, as Scala 3.0 doesn't seem to do the implicit conversion to Long anymore. + * This is not a bug, as Scala 2.13 already had a flag ("-Ywarn-implicit-widening") to turn on warnings/errors + * for that. + * + * The implicit conversion from Int to Long here doesn't lose precision and keeps backwards source compatibliity + * with previous releases. + */ + implicit def richDurationFromInt(numNanos: Int): RichDuration = + new RichDuration(numNanos.toLong) + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/MapOps.scala b/util-core/src/main/scala/com/twitter/conversions/MapOps.scala new file mode 100644 index 0000000000..fdf562cffb --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/MapOps.scala @@ -0,0 +1,140 @@ +package com.twitter.conversions + +import scala.collection.{SortedMap, immutable, mutable} + +/** + * Implicits for converting [[Map]]s. + * + * @example + * {{{ + * import com.twitter.conversions.MapOps._ + * + * Map(1 -> "a").mapKeys { _.toString } + * Map(1 -> "a").invert + * Map(1 -> "a", 2 -> "b").filterValues { _ == "b" } + * Map(2 -> "b", 1 -> "a").toSortedMap + * }}} + */ +object MapOps { + + implicit class RichMap[K, V](val self: Map[K, V]) extends AnyVal { + def mapKeys[T](func: K => T): Map[T, V] = MapOps.mapKeys(self, func) + + def invert: Map[V, Seq[K]] = MapOps.invert(self) + + def invertSingleValue: Map[V, K] = MapOps.invertSingleValue(self) + + def filterValues(func: V => Boolean): Map[K, V] = + MapOps.filterValues(self, func) + + def filterNotValues(func: V => Boolean): Map[K, V] = + MapOps.filterNotValues(self, func) + + def filterNotKeys(func: K => Boolean): Map[K, V] = + MapOps.filterNotKeys(self, func) + + def toSortedMap(implicit ordering: Ordering[K]): SortedMap[K, V] = + MapOps.toSortedMap(self) + } + + /** + * Transforms the keys of the map according to the given `func`. + * + * @param inputMap the map to transform the keys + * @param func the function literal which will be applied to the + * keys of the map + * @return the map with transformed keys and original values + */ + def mapKeys[K, V, T](inputMap: Map[K, V], func: K => T): Map[T, V] = { + for ((k, v) <- inputMap) yield { + func(k) -> v + } + } + + /** + * Inverts the map so that the input map's values are the distinct keys and + * the corresponding keys, represented in a sequence, as the values. + * + * @param inputMap the map to invert + * @return the inverted map + */ + def invert[K, V](inputMap: Map[K, V]): Map[V, Seq[K]] = { + val invertedMapWithBuilderValues = mutable.Map.empty[V, mutable.Builder[K, Seq[K]]] + for ((k, v) <- inputMap) { + val valueBuilder = invertedMapWithBuilderValues.getOrElseUpdate(v, Seq.newBuilder[K]) + valueBuilder += k + } + + val invertedMap = immutable.Map.newBuilder[V, Seq[K]] + for ((k, valueBuilder) <- invertedMapWithBuilderValues) { + invertedMap += (k -> valueBuilder.result()) + } + + invertedMap.result() + } + + /** + * Inverts the map so that every input map value becomes the key and the + * input map key becomes the value. + * + * @param inputMap the map to invert + * @return the inverted map + */ + def invertSingleValue[K, V](inputMap: Map[K, V]): Map[V, K] = { + inputMap map { _.swap } + } + + /** + * Filters the pairs in the map which values satisfy the predicate + * represented by `func`. + * + * @param inputMap the map which will be filtered + * @param func the predicate that needs to be satisfied to + * select a key-value pair + * @return the filtered map + */ + def filterValues[K, V](inputMap: Map[K, V], func: V => Boolean): Map[K, V] = { + inputMap filter { + case (_, value) => + func(value) + } + } + + /** + * Filters the pairs in the map which values do NOT satisfy the + * predicate represented by `func`. + * + * @param inputMap the map which will be filtered + * @param func the predicate that needs to be satisfied to NOT + * select a key-value pair + * @return the filtered map + */ + def filterNotValues[K, V](inputMap: Map[K, V], func: V => Boolean): Map[K, V] = + filterValues(inputMap, (v: V) => !func(v)) + + /** + * Filters the pairs in the map which keys do NOT satisfy the + * predicate represented by `func`. + * + * @param inputMap the map which will be filtered + * @param func the predicate that needs to be satisfied to NOT + * select a key-value pair + * @return the filtered map + */ + def filterNotKeys[K, V](inputMap: Map[K, V], func: K => Boolean): Map[K, V] = { + // use filter instead of filterKeys(deprecated since 2.13.0) for cross-building. + inputMap.filter { case (key, _) => !func(key) } + } + + /** + * Sorts the map by the keys and returns a SortedMap. + * + * @param inputMap the map to sort + * @param ordering the order in which to sort the map + * @return a SortedMap + */ + def toSortedMap[K, V](inputMap: Map[K, V])(implicit ordering: Ordering[K]): SortedMap[K, V] = { + SortedMap[K, V](inputMap.toSeq: _*) + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/PercentOps.scala b/util-core/src/main/scala/com/twitter/conversions/PercentOps.scala new file mode 100644 index 0000000000..ea2484726d --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/PercentOps.scala @@ -0,0 +1,33 @@ +package com.twitter.conversions + +/** + * Implicits for turning `x.percent` (where `x` is an `Int` or `Double`) into a `Double` + * scaled to where 1.0 is 100 percent. + * + * @note Negative values, fractional values, and values greater than 100 are permitted. + * + * @example + * {{{ + * import com.twitter.conversions.PercentOps._ + * + * 1.percent == 0.01 + * 100.percent == 1.0 + * 99.9.percent == 0.999 + * 500.percent == 5.0 + * -10.percent == -0.1 + * }}} + */ +object PercentOps { + + private val BigDecimal100 = BigDecimal(100.0) + + implicit class RichPercent(val value: Double) extends AnyVal { + // convert wrapped to BigDecimal to preserve precision when dividing Doubles + def percent: Double = + if (value.equals(Double.NaN) + || value.equals(Double.PositiveInfinity) + || value.equals(Double.NegativeInfinity)) value + else (BigDecimal(value) / BigDecimal100).doubleValue + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/SeqOps.scala b/util-core/src/main/scala/com/twitter/conversions/SeqOps.scala new file mode 100644 index 0000000000..d9f337d8ea --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/SeqOps.scala @@ -0,0 +1,117 @@ +package com.twitter.conversions + +import scala.annotation.tailrec + +abstract class SeqUtilCompanion { + def hasKnownSize[A](self: Seq[A]): Boolean +} + +/** + * Implicits for converting [[Seq]]s. + * + * @example + * {{{ + * import com.twitter.conversions.SeqOps._ + * + * Seq("a", "and").createMap(_.size, _) + * Seq("a", "and").groupBySingleValue(_.size) + * Seq("a", "and").findItemAfter("a") + * + * var numStrings = 0 + * Seq(1, 2) foreachPartial { + * case e: Int => + * numStrings += 1 + * } + * }}} + */ +object SeqOps { + + implicit class RichSeq[A](val self: Seq[A]) extends AnyVal { + + def createMap[K, V](keys: A => K, values: A => V): Map[K, V] = + SeqOps.createMap(self, keys, values) + + def createMap[K, V](values: A => V): Map[A, V] = + SeqOps.createMap(self, values) + + def foreachPartial(pf: PartialFunction[A, Unit]): Unit = + SeqOps.foreachPartial(self, pf) + + /** + * Creates a map from the given sequence by mapping each transformed + * element (the key) by applying the function `keys` to the sequence and + * using the element in the sequence as the value. Chooses the last + * element in the seq when a key collision occurs. + * + * @param keys the function that transforms a given element in the seq to a key + * @return a map of the given sequence + */ + def groupBySingleValue[B](keys: A => B): Map[B, A] = + createMap(keys, identity) + + def findItemAfter(itemToFind: A): Option[A] = + SeqOps.findItemAfter(self, itemToFind) + } + + /** + * Creates a map from the given sequence by applying the functions `keys` + * and `values` to the sequence. + * + * @param sequence the seq to convert to a map + * @param keys the function that transforms a given element in the seq to a key + * @param values the function that transforms a given element in the seq to a value + * @return a map of the given sequence + */ + def createMap[A, K, V](sequence: Seq[A], keys: A => K, values: A => V): Map[K, V] = { + sequence.iterator.map { elem => + keys(elem) -> values(elem) + }.toMap + } + + /** + * Creates a map from the given sequence by mapping each element of the + * sequence (the key) to the transformed element (the value) by applying + * the function `values`. + * + * @param sequence the seq to convert to a map + * @param values the function that transforms a given element in the seq to a value + * @return a map of the given sequence + */ + def createMap[A, K, V](sequence: Seq[A], values: A => V): Map[A, V] = { + sequence.iterator.map { elem => + elem -> values(elem) + }.toMap + } + + /** + * Applies the given partial function to the sequence. + * + * @param sequence the seq onto which the partial function is applied + * @param pf the partial function to apply to the seq + */ + def foreachPartial[A](sequence: Seq[A], pf: PartialFunction[A, Unit]): Unit = { + sequence.foreach { elem => + if (pf.isDefinedAt(elem)) { + pf(elem) + } + } + } + + /** + * Finds the element after the given element to find. + * + * @param sequence the seq to scan for the element + * @param itemToFind the element before the one returned + * @return an Option of the element after the item to find + */ + def findItemAfter[A](sequence: Seq[A], itemToFind: A): Option[A] = { + @tailrec + def recurse(itemToFind: A, seq: Seq[A]): Option[A] = seq match { + case Seq(x, xs @ _*) if x == itemToFind => xs.headOption + case Seq(x, xs @ _*) => recurse(itemToFind, xs) + case Seq() => None + } + recurse(itemToFind, sequence) + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/StorageUnitOps.scala b/util-core/src/main/scala/com/twitter/conversions/StorageUnitOps.scala new file mode 100644 index 0000000000..c7eac3846f --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/StorageUnitOps.scala @@ -0,0 +1,50 @@ +package com.twitter.conversions + +import com.twitter.util.StorageUnit + +import scala.language.implicitConversions + +/** + * Implicits for writing readable [[com.twitter.util.StorageUnit]]s. + * + * @example + * {{{ + * import com.twitter.conversions.StorageUnitOps._ + * + * 5.bytes + * 1.kilobyte + * 256.gigabytes + * }}} + */ +object StorageUnitOps { + + implicit class RichStorageUnit(val numBytes: Long) extends AnyVal { + def byte: StorageUnit = bytes + def bytes: StorageUnit = StorageUnit.fromBytes(numBytes) + def kilobyte: StorageUnit = kilobytes + def kilobytes: StorageUnit = StorageUnit.fromKilobytes(numBytes) + def megabyte: StorageUnit = megabytes + def megabytes: StorageUnit = StorageUnit.fromMegabytes(numBytes) + def gigabyte: StorageUnit = gigabytes + def gigabytes: StorageUnit = StorageUnit.fromGigabytes(numBytes) + def terabyte: StorageUnit = terabytes + def terabytes: StorageUnit = StorageUnit.fromTerabytes(numBytes) + def petabyte: StorageUnit = petabytes + def petabytes: StorageUnit = StorageUnit.fromPetabytes(numBytes) + + def thousand: Long = numBytes * 1000 + def million: Long = numBytes * 1000 * 1000 + def billion: Long = numBytes * 1000 * 1000 * 1000 + } + + /** + * Forwarder for Int, as Scala 3.0 doesn't seem to do the implicit conversion to Long anymore. + * This is not a bug, as Scala 2.13 already had a flag ("-Ywarn-implicit-widening") to turn on warnings/errors + * for that. + * + * The implicit conversion from Int to Long here doesn't lose precision and keeps backwards source compatibliity + * with previous releases. + */ + implicit def richStorageUnitFromInt(numBytes: Int): RichStorageUnit = + new RichStorageUnit(numBytes.toLong) +} diff --git a/util-core/src/main/scala/com/twitter/conversions/StringOps.scala b/util-core/src/main/scala/com/twitter/conversions/StringOps.scala new file mode 100644 index 0000000000..9e5ad18f02 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/StringOps.scala @@ -0,0 +1,314 @@ +package com.twitter.conversions + +import scala.util.matching.Regex + +object StringOps { + + // we intentionally don't unquote "\$" here, so it can be used to escape interpolation later. + private val UNQUOTE_RE = """\\(u[\dA-Fa-f]{4}|x[\dA-Fa-f]{2}|[/rnt\"\\])""".r + + private val QUOTE_RE = "[\u0000-\u001f\u007f-\uffff\\\\\"]".r + + implicit final class RichString(val string: String) extends AnyVal { + + /** + * For every section of a string that matches a regular expression, call + * a function to determine a replacement (as in python's + * `re.sub`). The function will be passed the Matcher object + * corresponding to the substring that matches the pattern, and that + * substring will be replaced by the function's result. + * + * For example, this call: + * {{{ + * "ohio".regexSub("""h.""".r) { m => "n" } + * }}} + * will return the string `"ono"`. + * + * The matches are found using `Matcher.find()` and so + * will obey all the normal java rules (the matches will not overlap, + * etc). + * + * @param re the regex pattern to replace + * @param replace a function that takes Regex.MatchData objects and + * returns a string to substitute + * @return the resulting string with replacements made + */ + def regexSub(re: Regex)(replace: Regex.MatchData => String): String = { + var offset = 0 + val out = new StringBuilder() + + for (m <- re.findAllIn(string).matchData) { + if (m.start > offset) { + out.append(string.substring(offset, m.start)) + } + + out.append(replace(m)) + offset = m.end + } + + if (offset < string.length) { + out.append(string.substring(offset)) + } + out.toString + } + + /** + * Quote a string so that unprintable chars (in ASCII) are represented by + * C-style backslash expressions. For example, a raw linefeed will be + * translated into `"\n"`. Control codes (anything below 0x20) + * and unprintables (anything above 0x7E) are turned into either + * `"\xHH"` or `"\\uHHHH"` expressions, depending on + * their range. Embedded backslashes and double-quotes are also quoted. + * + * @return a quoted string, suitable for ASCII display + */ + def quoteC(): String = { + regexSub(QUOTE_RE) { m => + m.matched.charAt(0) match { + case '\r' => "\\r" + case '\n' => "\\n" + case '\t' => "\\t" + case '"' => "\\\"" + case '\\' => "\\\\" + case c => + if (c <= 255) { + "\\x%02x".format(c.asInstanceOf[Int]) + } else { + "\\u%04x" format c.asInstanceOf[Int] + } + } + } + } + + /** + * Unquote an ASCII string that has been quoted in a style like + * `quoteC` and convert it into a standard unicode string. + * `"\\uHHHH"` and `"\xHH"` expressions are unpacked + * into unicode characters, as well as `"\r"`, `"\n"`, + * `"\t"`, `"\\"`, and `'\"'`. + * + * @return an unquoted unicode string + */ + def unquoteC(): String = { + regexSub(UNQUOTE_RE) { m => + val ch = m.group(1).charAt(0) match { + // holy crap! this is terrible: + case 'u' => + Character.valueOf(Integer.valueOf(m.group(1).substring(1), 16).asInstanceOf[Int].toChar) + case 'x' => + Character.valueOf(Integer.valueOf(m.group(1).substring(1), 16).asInstanceOf[Int].toChar) + case 'r' => '\r' + case 'n' => '\n' + case 't' => '\t' + case x => x + } + ch.toString + } + } + + /** + * Turn a string of hex digits into a byte array. This does the exact + * opposite of `Array[Byte]#hexlify`. + */ + def unhexlify(): Array[Byte] = { + val buffer = new Array[Byte]((string.length + 1) / 2) + string.grouped(2).toSeq.zipWithIndex.foreach { + case (substr, i) => + buffer(i) = Integer.parseInt(substr, 16).toByte + } + buffer + } + + /** + * Converts foo_bar to fooBar (first letter lowercased) + * + * For example: + * + * {{{ + * "hello_world".toCamelCase + * }}} + * + * will return the String `"helloWorld"`. + */ + def toCamelCase: String = StringOps.toCamelCase(string) + + /** + * Converts foo_bar to FooBar (first letter uppercased) + * + * For example: + * + * {{{ + * "hello_world".toPascalCase + * }}} + * + * will return the String `"HelloWorld"`. + */ + def toPascalCase: String = StringOps.toPascalCase(string) + + /** + * Converts FooBar to foo_bar (all lowercase) + * + * For example: + * + * {{{ + * "HelloWorld".toSnakeCase + * }}} + * + * will return the String `"hello_world"`. + */ + def toSnakeCase: String = StringOps.toSnakeCase(string) + + /** + * Wrap the string in an Option, returning None when it is null or empty. + */ + def toOption: Option[String] = StringOps.toOption(string) + + /** + * Return the string unless it is null or empty, in which + * case return the default. + * + * @param default the string to return if inputString is null or empty + * @return either the inputString or the default + */ + def getOrElse(default: => String): String = StringOps.getOrElse(string, default) + } + + def hexlify(array: Array[Byte], from: Int, to: Int): String = { + val out = new StringBuffer + for (i <- from until to) { + val b = array(i) + val s = (b.toInt & 0xff).toHexString + if (s.length < 2) { + out.append('0') + } + out.append(s) + } + out.toString + } + + private[this] val SnakeCaseRegexFirstPass = """([A-Z]+)([A-Z][a-z])""".r + private[this] val SnakeCaseRegexSecondPass = """([a-z\d])([A-Z])""".r + + /** + * Turn a string of format "FooBar" into snake case "foo_bar" + * + * @note toSnakeCase is not reversible, ie. in general the following will _not_ be true: + * + * {{{ + * s == toCamelCase(toSnakeCase(s)) + * }}} + * + * Copied from the Lift Framework and RENAMED: + * https://github.com/lift/framework/blob/master/core/util/src/main/scala/net/liftweb/util/StringHelpers.scala + * Apache 2.0 License: https://github.com/lift/framework/blob/master/LICENSE.txt + * + * @return the underscored string + */ + def toSnakeCase(name: String): String = + SnakeCaseRegexSecondPass + .replaceAllIn( + SnakeCaseRegexFirstPass + .replaceAllIn(name, "$1_$2"), + "$1_$2" + ) + .toLowerCase + + /** + * Turns a string of format "foo_bar" into PascalCase "FooBar" + * + * @note Camel case may start with a capital letter (called "Pascal Case" or "Upper Camel Case") or, + * especially often in programming languages, with a lowercase letter. In this code's + * perspective, PascalCase means the first char is capitalized while camelCase means the first + * char is lowercased. In general both can be considered equivalent although by definition + * "CamelCase" is a valid camel-cased word. Hence, PascalCase can be considered to be a + * subset of camelCase. + * + * Copied from the "lift" framework and RENAMED: + * https://github.com/lift/framework/blob/master/core/util/src/main/scala/net/liftweb/util/StringHelpers.scala + * Apache 2.0 License: https://github.com/lift/framework/blob/master/LICENSE.txt + * Functional code courtesy of Jamie Webb (j@jmawebb.cjb.net) 2006/11/28 + * + * @param name the String to PascalCase + * @return the PascalCased string + */ + def toPascalCase(name: String): String = { + def loop(x: List[Char]): List[Char] = (x: @unchecked) match { + case '_' :: '_' :: rest => loop('_' :: rest) + case '_' :: c :: rest => Character.toUpperCase(c) :: loop(rest) + case '_' :: Nil => Nil + case c :: rest => c :: loop(rest) + case Nil => Nil + } + + if (name == null) { + "" + } else { + loop('_' :: name.toList).mkString + } + } + + /** + * Turn a string of format "foo_bar" into camelCase with the first letter in lower case: "fooBar" + * + * This function is especially used to camelCase method names. + * + * @note Camel case may start with a capital letter (called "Pascal Case" or "Upper Camel Case") or, + * especially often in programming languages, with a lowercase letter. In this code's + * perspective, PascalCase means the first char is capitalized while camelCase means the first + * char is lowercased. In general both can be considered equivalent although by definition + * "CamelCase" is a valid camel-cased word. Hence, PascalCase can be considered to be a + * subset of camelCase. + * + * Copied from the Lift Framework and RENAMED: + * https://github.com/lift/framework/blob/master/core/util/src/main/scala/net/liftweb/util/StringHelpers.scala + * Apache 2.0 License: https://github.com/lift/framework/blob/master/LICENSE.txt + * + * @param name the String to camelCase + * @return the camelCased string + */ + def toCamelCase(name: String): String = { + val tmp: String = toPascalCase(name) + if (tmp.length == 0) { + "" + } else { + tmp.substring(0, 1).toLowerCase + tmp.substring(1) + } + } + + /** + * Wrap the string in an Option, returning None when it is null or empty. + * + * @param inputString the string to wrap in an Option + * @return an Option of the given string + */ + def toOption(inputString: String): Option[String] = { + if (inputString == null || inputString.isEmpty) + None + else + Some(inputString) + } + + /** + * Return the string unless it is null or empty, in which + * case return the default. + * + * @param inputString the string to "get" + * @param default the string to return if inputString is null or empty + * @return either the inputString or the default + */ + def getOrElse(inputString: String, default: => String): String = { + if (inputString == null || inputString.isEmpty) + default + else + inputString + } + + implicit final class RichByteArray(val bytes: Array[Byte]) extends AnyVal { + + /** + * Turn an `Array[Byte]` into a string of hex digits. + */ + def hexlify: String = StringOps.hexlify(bytes, 0, bytes.length) + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/ThreadOps.scala b/util-core/src/main/scala/com/twitter/conversions/ThreadOps.scala new file mode 100644 index 0000000000..a6589c4c78 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/ThreadOps.scala @@ -0,0 +1,19 @@ +package com.twitter.conversions + +import java.util.concurrent.Callable +import scala.language.implicitConversions + +/** + * Implicits for turning a block of code into a Runnable or Callable. + */ +object ThreadOps { + + implicit def makeRunnable(f: => Unit): Runnable = new Runnable() { + def run(): Unit = f + } + + implicit def makeCallable[T](f: => T): Callable[T] = new Callable[T]() { + def call(): T = f + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/TupleOps.scala b/util-core/src/main/scala/com/twitter/conversions/TupleOps.scala new file mode 100644 index 0000000000..3669681599 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/TupleOps.scala @@ -0,0 +1,172 @@ +package com.twitter.conversions + +import scala.collection.{SortedMap, immutable, mutable} +import scala.math.Ordering + +/** + * Implicits for converting tuples. + * + * @example + * {{{ + * import com.twitter.conversions.TupleOps._ + * + * Seq(1 -> "tada", 2 -> "lala").toKeys + * Seq(1 -> "tada", 2 -> "lala").mapValues(_.size) + * Seq(1 -> "tada", 2 -> "lala").groupByKeyAndReduce(_ + "," + _) + * Seq(1 -> "tada", 2 -> "lala").sortByKey + * }}} + */ +object TupleOps { + + implicit class RichTuples[A, B](val self: Iterable[(A, B)]) extends AnyVal { + def toKeys: Seq[A] = TupleOps.toKeys(self) + + def toKeySet: Set[A] = TupleOps.toKeySet(self) + + def toValues: Seq[B] = TupleOps.toValues(self) + + def mapValues[C](func: B => C): Seq[(A, C)] = + TupleOps.mapValues(self, func) + + def groupByKey: Map[A, Seq[B]] = + TupleOps.groupByKey(self) + + def groupByKeyAndReduce(reduceFunc: (B, B) => B): Map[A, B] = + TupleOps.groupByKeyAndReduce(self, reduceFunc) + + def sortByKey(implicit ord: Ordering[A]): Seq[(A, B)] = + TupleOps.sortByKey(self) + + def toSortedMap(implicit ord: Ordering[A]): SortedMap[A, B] = + TupleOps.toSortedMap(self) + } + + /** + * Returns a new seq of the first element of each tuple. + * + * @param tuples the tuples for which all the first elements will be returned + * @return a seq of the first element of each tuple + */ + def toKeys[A, B](tuples: Iterable[(A, B)]): Seq[A] = { + tuples.toSeq map { case (key, value) => key } + } + + /** + * Returns a new set of the first element of each tuple. + * + * @param tuples the tuples for which all the first elements will be returned + * @return a set of the first element of each tuple + */ + def toKeySet[A, B](tuples: Iterable[(A, B)]): Set[A] = { + toKeys(tuples).toSet + } + + /** + * Returns a new seq of the second element of each tuple. + * + * @param tuples the tuples for which all the second elements + * will be returned + * @return a seq of the second element of each tuple + */ + def toValues[A, B](tuples: Iterable[(A, B)]): Seq[B] = { + tuples.toSeq map { case (key, value) => value } + } + + /** + * Applies `func` to the second element of each tuple. + * + * @param tuples the tuples to which the func is applied + * @param func the function literal which will be applied + * to the second element of each tuple + * @return a seq of `func` transformed tuples + */ + def mapValues[A, B, C](tuples: Iterable[(A, B)], func: B => C): Seq[(A, C)] = { + tuples.toSeq map { + case (key, value) => + key -> func(value) + } + } + + /** + * Groups the tuples by matching the first element (key) in each tuple + * returning a map from the key to a seq of the according second elements. + * + * For example: + * + * {{{ + * groupByKey(Seq(1 -> "a", 1 -> "b", 2 -> "ab")) + * }}} + * + * will return Map(2 -> Seq("ab"), 1 -> Seq("a", "b")) + * + * @param tuples the tuples or organize into a grouped map + * @return a map of matching first elements to a sequence of + * their second elements + */ + def groupByKey[A, B](tuples: Iterable[(A, B)]): Map[A, Seq[B]] = { + val mutableMapBuilder = mutable.Map.empty[A, mutable.Builder[B, Seq[B]]] + for ((a, b) <- tuples) { + val seqBuilder = mutableMapBuilder.getOrElseUpdate(a, immutable.Seq.newBuilder[B]) + seqBuilder += b + } + + val mapBuilder = immutable.Map.newBuilder[A, Seq[B]] + for ((k, v) <- mutableMapBuilder) { + mapBuilder += ((k, v.result())) + } + + mapBuilder.result() + } + + /** + * First groups the tuples by matching first elements creating a + * map of the first elements to a seq of the according second elements, + * then reduces the seq to a single value according to the + * applied `reduceFunc`. + * + * For example: + * + * {{{ + * groupByKeyAndReduce(Seq(1 -> 5, 1 -> 6, 2 -> 7), _ + _) + * }}} + * + * will return Map(2 -> 7, 1 -> 11) + * + * @param tuples the tuples or organize into a grouped map + * @param reduceFunc the function to apply to the seq of second + * elements that have been grouped together + * @return a map of matching first elements to a value after + * reducing the grouped values with the `reduceFunc` + */ + def groupByKeyAndReduce[A, B](tuples: Iterable[(A, B)], reduceFunc: (B, B) => B): Map[A, B] = { + // use map instead of mapValues(deprecated since 2.13.0) for cross-building. + groupByKey(tuples).map { + case (k, values) => + k -> values.reduce(reduceFunc) + } + } + + /** + * Sorts the tuples by the first element in each tuple. + * + * @param tuples the tuples to sort + * @param ord the order in which to sort the tuples + * @return the sorted tuples + */ + def sortByKey[A, B](tuples: Iterable[(A, B)])(implicit ord: Ordering[A]): Seq[(A, B)] = { + tuples.toSeq sortBy { case (key, value) => key } + } + + /** + * Sorts the tuples by the first element in each tuple and + * returns a SortedMap. + * + * @param tuples the tuples to sort + * @param ord the order in which to sort the tuples + * @return a SortedMap of the tuples + */ + def toSortedMap[A, B](tuples: Iterable[(A, B)])(implicit ord: Ordering[A]): SortedMap[A, B] = { + SortedMap(tuples.toSeq: _*) + } + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/U64Ops.scala b/util-core/src/main/scala/com/twitter/conversions/U64Ops.scala new file mode 100644 index 0000000000..8ff3586440 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/conversions/U64Ops.scala @@ -0,0 +1,60 @@ +package com.twitter.conversions + +import scala.language.implicitConversions + +object U64Ops { + + /** + * Parses this HEX string as an unsigned 64-bit long value. Be careful, this can throw + * `NumberFormatException`. + * + * See `java.lang.Long.parseUnsignedLong()` for more information. + */ + implicit class RichU64String(val self: String) extends AnyVal { + def toU64Long: Long = java.lang.Long.parseUnsignedLong(self, 16) + } + + private[this] val lookupTable: Array[Char] = + (('0' to '9') ++ ('a' to 'f')).toArray + + private[this] def hex(l: Long): String = { + val chs = new Array[Char](16) + chs(0) = lookupTable((l >> 60 & 0xf).toInt) + chs(1) = lookupTable((l >> 56 & 0xf).toInt) + chs(2) = lookupTable((l >> 52 & 0xf).toInt) + chs(3) = lookupTable((l >> 48 & 0xf).toInt) + chs(4) = lookupTable((l >> 44 & 0xf).toInt) + chs(5) = lookupTable((l >> 40 & 0xf).toInt) + chs(6) = lookupTable((l >> 36 & 0xf).toInt) + chs(7) = lookupTable((l >> 32 & 0xf).toInt) + chs(8) = lookupTable((l >> 28 & 0xf).toInt) + chs(9) = lookupTable((l >> 24 & 0xf).toInt) + chs(10) = lookupTable((l >> 20 & 0xf).toInt) + chs(11) = lookupTable((l >> 16 & 0xf).toInt) + chs(12) = lookupTable((l >> 12 & 0xf).toInt) + chs(13) = lookupTable((l >> 8 & 0xf).toInt) + chs(14) = lookupTable((l >> 4 & 0xf).toInt) + chs(15) = lookupTable((l & 0xf).toInt) + new String(chs, 0, 16) + } + + /** + * Converts this unsigned 64-bit long value into a 16-character HEX string. + */ + implicit class RichU64Long(val self: Long) extends AnyVal { + + def toU64HexString: String = hex(self) + } + + /** + * Forwarder for Int, as Scala 3.0 doesn't seem to do the implicit conversion to Long anymore. + * This is not a bug, as Scala 2.13 already had a flag ("-Ywarn-implicit-widening") to turn on warnings/errors + * for that. + * + * The implicit conversion from Int to Long here doesn't lose precision and keeps backwards source compatibliity + * with previous releases. + */ + implicit def richU64LongFromInt(self: Int): RichU64Long = + new RichU64Long(self.toLong) + +} diff --git a/util-core/src/main/scala/com/twitter/conversions/percent.scala b/util-core/src/main/scala/com/twitter/conversions/percent.scala deleted file mode 100644 index e49fefb8b8..0000000000 --- a/util-core/src/main/scala/com/twitter/conversions/percent.scala +++ /dev/null @@ -1,37 +0,0 @@ -package com.twitter.conversions - -import scala.language.implicitConversions - -/** - * Implicits for turning x.percent (where x is an Int or Double) into a Double scaled to - * where 1.0 is 100 percent. - * - * @note Negative values, fractional values, and values greater than 100 are permitted. - * - * @example - * {{{ - * 1.percent == 0.01 - * 100.percent == 1.0 - * 99.9.percent == 0.999 - * 500.percent == 5.0 - * -10.percent == -0.1 - * }}} - */ -object percent { - - private val BigDecimal100 = BigDecimal(100.0) - - class RichDoublePercent private[conversions](val wrapped: Double) extends AnyVal { - // convert wrapped to BigDecimal to preserve precision when dividing Doubles - def percent: Double = - if (wrapped.equals(Double.NaN) - || wrapped.equals(Double.PositiveInfinity) - || wrapped.equals(Double.NegativeInfinity)) wrapped - else (BigDecimal(wrapped) / BigDecimal100).doubleValue - } - - implicit def intPercentToFractionalDouble(i: Int): RichDoublePercent = - new RichDoublePercent(i) - implicit def doublePercentToFractionalDouble(d: Double): RichDoublePercent = - new RichDoublePercent(d) -} diff --git a/util-core/src/main/scala/com/twitter/conversions/storage.scala b/util-core/src/main/scala/com/twitter/conversions/storage.scala deleted file mode 100644 index a7dfef97e6..0000000000 --- a/util-core/src/main/scala/com/twitter/conversions/storage.scala +++ /dev/null @@ -1,44 +0,0 @@ -/* - * Copyright 2010 Twitter Inc. - * - * 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 com.twitter.conversions - -import com.twitter.util.StorageUnit -import scala.language.implicitConversions - -object storage { - class RichWholeNumber(wrapped: Long) { - def byte: StorageUnit = bytes - def bytes: StorageUnit = StorageUnit.fromBytes(wrapped) - def kilobyte: StorageUnit = kilobytes - def kilobytes: StorageUnit = StorageUnit.fromKilobytes(wrapped) - def megabyte: StorageUnit = megabytes - def megabytes: StorageUnit = StorageUnit.fromMegabytes(wrapped) - def gigabyte: StorageUnit = gigabytes - def gigabytes: StorageUnit = StorageUnit.fromGigabytes(wrapped) - def terabyte: StorageUnit = terabytes - def terabytes: StorageUnit = StorageUnit.fromTerabytes(wrapped) - def petabyte: StorageUnit = petabytes - def petabytes: StorageUnit = StorageUnit.fromPetabytes(wrapped) - - def thousand: Long = wrapped * 1000 - def million: Long = wrapped * 1000 * 1000 - def billion: Long = wrapped * 1000 * 1000 * 1000 - } - - implicit def intToStorageUnitableWholeNumber(i: Int): RichWholeNumber = new RichWholeNumber(i) - implicit def longToStorageUnitableWholeNumber(l: Long): RichWholeNumber = new RichWholeNumber(l) -} diff --git a/util-core/src/main/scala/com/twitter/conversions/string.scala b/util-core/src/main/scala/com/twitter/conversions/string.scala deleted file mode 100644 index aa6aae5000..0000000000 --- a/util-core/src/main/scala/com/twitter/conversions/string.scala +++ /dev/null @@ -1,162 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.conversions - -import scala.language.implicitConversions -import scala.util.matching.Regex - -object string { - final class RichString(wrapped: String) { - - /** - * For every section of a string that matches a regular expression, call - * a function to determine a replacement (as in python's - * `re.sub`). The function will be passed the Matcher object - * corresponding to the substring that matches the pattern, and that - * substring will be replaced by the function's result. - * - * For example, this call: - * {{{ - * "ohio".regexSub("""h.""".r) { m => "n" } - * }}} - * will return the string `"ono"`. - * - * The matches are found using `Matcher.find()` and so - * will obey all the normal java rules (the matches will not overlap, - * etc). - * - * @param re the regex pattern to replace - * @param replace a function that takes Regex.MatchData objects and - * returns a string to substitute - * @return the resulting string with replacements made - */ - def regexSub(re: Regex)(replace: Regex.MatchData => String): String = { - var offset = 0 - val out = new StringBuilder() - - for (m <- re.findAllIn(wrapped).matchData) { - if (m.start > offset) { - out.append(wrapped.substring(offset, m.start)) - } - - out.append(replace(m)) - offset = m.end - } - - if (offset < wrapped.length) { - out.append(wrapped.substring(offset)) - } - out.toString - } - - private val QUOTE_RE = "[\u0000-\u001f\u007f-\uffff\\\\\"]".r - - /** - * Quote a string so that unprintable chars (in ASCII) are represented by - * C-style backslash expressions. For example, a raw linefeed will be - * translated into `"\n"`. Control codes (anything below 0x20) - * and unprintables (anything above 0x7E) are turned into either - * `"\xHH"` or `"\\uHHHH"` expressions, depending on - * their range. Embedded backslashes and double-quotes are also quoted. - * - * @return a quoted string, suitable for ASCII display - */ - def quoteC(): String = { - regexSub(QUOTE_RE) { m => - m.matched.charAt(0) match { - case '\r' => "\\r" - case '\n' => "\\n" - case '\t' => "\\t" - case '"' => "\\\"" - case '\\' => "\\\\" - case c => - if (c <= 255) { - "\\x%02x".format(c.asInstanceOf[Int]) - } else { - "\\u%04x" format c.asInstanceOf[Int] - } - } - } - } - - // we intentionally don't unquote "\$" here, so it can be used to escape interpolation later. - private val UNQUOTE_RE = """\\(u[\dA-Fa-f]{4}|x[\dA-Fa-f]{2}|[/rnt\"\\])""".r - - /** - * Unquote an ASCII string that has been quoted in a style like - * [[quoteC()]] and convert it into a standard unicode string. - * `"\\uHHHH"` and `"\xHH"` expressions are unpacked - * into unicode characters, as well as `"\r"`, `"\n"`, - * `"\t"`, `"\\"`, and `'\"'`. - * - * @return an unquoted unicode string - */ - def unquoteC() = { - regexSub(UNQUOTE_RE) { m => - val ch = m.group(1).charAt(0) match { - // holy crap! this is terrible: - case 'u' => - Character.valueOf(Integer.valueOf(m.group(1).substring(1), 16).asInstanceOf[Int].toChar) - case 'x' => - Character.valueOf(Integer.valueOf(m.group(1).substring(1), 16).asInstanceOf[Int].toChar) - case 'r' => '\r' - case 'n' => '\n' - case 't' => '\t' - case x => x - } - ch.toString - } - } - - /** - * Turn a string of hex digits into a byte array. This does the exact - * opposite of `Array[Byte]#hexlify`. - */ - def unhexlify(): Array[Byte] = { - val buffer = new Array[Byte]((wrapped.length + 1) / 2) - wrapped.grouped(2).toSeq.zipWithIndex.foreach { - case (substr, i) => - buffer(i) = Integer.parseInt(substr, 16).toByte - } - buffer - } - } - - def hexlify(array: Array[Byte], from: Int, to: Int): String = { - val out = new StringBuffer - for (i <- from until to) { - val b = array(i) - val s = (b.toInt & 0xff).toHexString - if (s.length < 2) { - out append '0' - } - out append s - } - out.toString - } - - final class RichByteArray(wrapped: Array[Byte]) { - - /** - * Turn an `Array[Byte]` into a string of hex digits. - */ - def hexlify: String = string.hexlify(wrapped, 0, wrapped.length) - } - - implicit def stringToConfiggyString(s: String): RichString = new RichString(s) - implicit def byteArrayToConfiggyByteArray(b: Array[Byte]): RichByteArray = new RichByteArray(b) -} diff --git a/util-core/src/main/scala/com/twitter/conversions/thread.scala b/util-core/src/main/scala/com/twitter/conversions/thread.scala deleted file mode 100644 index ff4afab986..0000000000 --- a/util-core/src/main/scala/com/twitter/conversions/thread.scala +++ /dev/null @@ -1,29 +0,0 @@ -/* - * Copyright 2010 Twitter Inc. - * - * 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 com.twitter.conversions - -import java.util.concurrent.Callable -import scala.language.implicitConversions - -/** - * Implicits for turning a block of code into a Runnable or Callable. - */ -object thread { - implicit def makeRunnable(f: => Unit): Runnable = new Runnable() { def run(): Unit = f } - - implicit def makeCallable[T](f: => T): Callable[T] = new Callable[T]() { def call(): T = f } -} diff --git a/util-core/src/main/scala/com/twitter/conversions/time.scala b/util-core/src/main/scala/com/twitter/conversions/time.scala deleted file mode 100644 index 993dbf4d6f..0000000000 --- a/util-core/src/main/scala/com/twitter/conversions/time.scala +++ /dev/null @@ -1,56 +0,0 @@ -/* - * Copyright 2010 Twitter Inc. - * - * 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 com.twitter.conversions - -import com.twitter.util.Duration -import java.util.concurrent.TimeUnit -import scala.language.implicitConversions - -object time { - class RichWholeNumber(wrapped: Long) { - def nanoseconds: Duration = Duration(wrapped, TimeUnit.NANOSECONDS) - def nanosecond: Duration = nanoseconds - def microseconds: Duration = Duration(wrapped, TimeUnit.MICROSECONDS) - def microsecond: Duration = microseconds - def milliseconds: Duration = Duration(wrapped, TimeUnit.MILLISECONDS) - def millisecond: Duration = milliseconds - def millis: Duration = milliseconds - def seconds: Duration = Duration(wrapped, TimeUnit.SECONDS) - def second: Duration = seconds - def minutes: Duration = Duration(wrapped, TimeUnit.MINUTES) - def minute: Duration = minutes - def hours: Duration = Duration(wrapped, TimeUnit.HOURS) - def hour: Duration = hours - def days: Duration = Duration(wrapped, TimeUnit.DAYS) - def day: Duration = days - } - - private val ZeroRichWholeNumber = new RichWholeNumber(0) { - override def nanoseconds: Duration = Duration.Zero - override def microseconds: Duration = Duration.Zero - override def milliseconds: Duration = Duration.Zero - override def seconds: Duration = Duration.Zero - override def minutes: Duration = Duration.Zero - override def hours: Duration = Duration.Zero - override def days: Duration = Duration.Zero - } - - implicit def intToTimeableNumber(i: Int): RichWholeNumber = - if (i == 0) ZeroRichWholeNumber else new RichWholeNumber(i) - implicit def longToTimeableNumber(l: Long): RichWholeNumber = - if (l == 0) ZeroRichWholeNumber else new RichWholeNumber(l) -} diff --git a/util-core/src/main/scala/com/twitter/io/ActivitySource.scala b/util-core/src/main/scala/com/twitter/io/ActivitySource.scala new file mode 100644 index 0000000000..4e3aefba8d --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/ActivitySource.scala @@ -0,0 +1,63 @@ +package com.twitter.io + +import com.twitter.util._ + +/** + * An ActivitySource provides access to observerable named variables. + */ +trait ActivitySource[+T] { + + /** + * Returns an [[com.twitter.util.Activity]] for a named T-typed variable. + */ + def get(name: String): Activity[T] + + /** + * Produces an ActivitySource which queries this ActivitySource, falling back to + * a secondary ActivitySource only when the primary result is ActivitySource.Failed. + * + * @param that the secondary ActivitySource + */ + def orElse[U >: T](that: ActivitySource[U]): ActivitySource[U] = + new ActivitySource.OrElse(this, that) +} + +object ActivitySource { + + /** + * A Singleton exception to indicate that an ActivitySource failed to find + * a named variable. + */ + object NotFound extends IllegalStateException + + /** + * An ActivitySource for observing file contents. Once observed, + * each file will be polled once per period. + */ + def forFiles( + period: Duration = Duration.fromSeconds(60) + )( + implicit timer: Timer + ): ActivitySource[Buf] = + new CachingActivitySource(new FilePollingActivitySource(period)(timer)) + + /** + * Create an ActivitySource for ClassLoader resources. + */ + def forClassLoaderResources( + cl: ClassLoader = ClassLoader.getSystemClassLoader + ): ActivitySource[Buf] = + new CachingActivitySource(new ClassLoaderActivitySource(cl)) + + private[ActivitySource] class OrElse[T, U >: T]( + primary: ActivitySource[T], + failover: ActivitySource[U]) + extends ActivitySource[U] { + def get(name: String): Activity[U] = { + primary.get(name) transform { + case Activity.Failed(_) => failover.get(name) + case state => Activity(Var.value(state)) + } + } + } +} diff --git a/util-core/src/main/scala/com/twitter/io/BUILD b/util-core/src/main/scala/com/twitter/io/BUILD new file mode 100644 index 0000000000..223da6fec4 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/BUILD @@ -0,0 +1,16 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-core-io", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + ], +) diff --git a/util-core/src/main/scala/com/twitter/io/Buf.scala b/util-core/src/main/scala/com/twitter/io/Buf.scala index 2d6869c204..6d59d3cb25 100644 --- a/util-core/src/main/scala/com/twitter/io/Buf.scala +++ b/util-core/src/main/scala/com/twitter/io/Buf.scala @@ -1,8 +1,10 @@ package com.twitter.io -import com.twitter.io.Buf.{IndexedThree, IndexedTwo} +import com.twitter.io.Buf.IndexedThree +import com.twitter.io.Buf.IndexedTwo import java.nio.ReadOnlyBufferException -import java.nio.charset.{Charset, StandardCharsets => JChar} +import java.nio.charset.Charset +import java.nio.charset.{StandardCharsets => JChar} import scala.annotation.switch import scala.collection.immutable.VectorBuilder @@ -13,7 +15,7 @@ import scala.collection.immutable.VectorBuilder * * @see [[com.twitter.io.Buf.ByteArray]] for an `Array[Byte]` backed * implementation. - * @see [[com.twitter.io.Buf.ByteBuffer]] for an `nio.ByteBuffer` backed + * @see [[com.twitter.io.Buf.ByteBuffer]] for a `java.nio.ByteBuffer` backed * implementation. * @see [[com.twitter.io.Buf.apply]] for creating a `Buf` from other `Bufs` * @see [[com.twitter.io.Buf.Empty]] for an empty `Buf`. @@ -37,9 +39,10 @@ abstract class Buf { outer => * want to manually set the first bytes of an `Array[Byte]` and then efficiently copy the * contents of this `Buf` to the remaining space. * - * @see [[write(ByteBuffer)]] for writing to nio buffers. + * @see [[com.twitter.io.Buf.write(output:java\.nio\.ByteBuffer):Unit* write(ByteBuffer)]] + * for writing to nio buffers. * - * @throws IllegalArgumentException when `output` is too small to + * @note Throws `IllegalArgumentException` when `output` is too small to * contain all the data. * * @note [[Buf]] implementors should use the helper [[checkWriteArgs]]. @@ -57,12 +60,13 @@ abstract class Buf { outer => * the data is destined for an IO operation, it may be preferable to provide a direct * nio `ByteBuffer` to ensure the avoidance of intermediate heap-based representations. * - * @see [[write(Array[Byte], Int)]] for writing to byte arrays. + * @see [[com.twitter.io.Buf.write(output:Array[Byte],off:Int):Unit* write(Array[Byte],Int)]] + * for writing to byte arrays. * - * @throws java.lang.IllegalArgumentException when `output` doesn't have enough + * @note Throws `java.lang.IllegalArgumentException` when `output` doesn't have enough * space as defined by `ByteBuffer.remaining()` to hold the contents of this `Buf`. * - * @throws ReadOnlyBufferException if the provided buffer is read-only. + * @note Throws `ReadOnlyBufferException` if the provided buffer is read-only. * * @note [[Buf]] implementors should use the helper [[checkWriteArgs]]. */ @@ -94,56 +98,60 @@ abstract class Buf { outer => else this match { // This could be much cleaner as a Tuple match, but we want to avoid the allocation. - case Buf.Composite.Impl(left: Vector[Buf], leftLength) => + case Buf.Composite(left: Vector[Buf], leftLength) => right match { - case Buf.Composite.Impl(right: Vector[Buf], rightLength) => - Buf.Composite.Impl(Buf.fastConcat(left, right), leftLength + rightLength) - case Buf.Composite.Impl(right: IndexedTwo, rightLength) => - Buf.Composite.Impl(left :+ right.a :+ right.b, leftLength + rightLength) - case Buf.Composite.Impl(right: IndexedThree, rightLength) => - Buf.Composite.Impl(left :+ right.a :+ right.b :+ right.c, leftLength + rightLength) + case Buf.Composite(right: Vector[Buf], rightLength) => + Buf.Composite(Buf.fastConcat(left, right), leftLength + rightLength) + case Buf.Composite(right: IndexedTwo, rightLength) => + Buf.Composite(left :+ right.a :+ right.b, leftLength + rightLength) + case Buf.Composite(right: IndexedThree, rightLength) => + Buf.Composite(left :+ right.a :+ right.b :+ right.c, leftLength + rightLength) case _ => - Buf.Composite.Impl(left :+ right, leftLength + right.length) + Buf.Composite(left :+ right, leftLength + right.length) } - case Buf.Composite.Impl(left: IndexedTwo, leftLength) => + case Buf.Composite(left: IndexedTwo, leftLength) => right match { - case Buf.Composite.Impl(right: Vector[Buf], rightLength) => - Buf.Composite.Impl(left.a +: left.b +: right, leftLength + rightLength) - case Buf.Composite.Impl(right: IndexedTwo, rightLength) => - Buf.Composite.Impl( + case Buf.Composite(right: Vector[Buf], rightLength) => + Buf.Composite(left.a +: left.b +: right, leftLength + rightLength) + case Buf.Composite(right: IndexedTwo, rightLength) => + Buf.Composite( (new VectorBuilder[Buf] += left.a += left.b += right.a += right.b).result(), leftLength + rightLength ) - case Buf.Composite.Impl(right: IndexedThree, rightLength) => - Buf.Composite.Impl( + case Buf.Composite(right: IndexedThree, rightLength) => + Buf.Composite( (new VectorBuilder[Buf] += left.a += left.b += right.a += right.b += right.c) .result(), leftLength + rightLength ) case _ => - Buf.Composite.Impl( - new IndexedThree(left.a, left.b, right), leftLength + right.length + Buf.Composite( + new IndexedThree(left.a, left.b, right), + leftLength + right.length ) } - case Buf.Composite.Impl(left: IndexedThree, leftLength) => + case Buf.Composite(left: IndexedThree, leftLength) => right match { - case Buf.Composite.Impl(right: Vector[Buf], rightLength) => - Buf.Composite.Impl(left.a +: left.b +: left.c +: right, leftLength + rightLength) - case Buf.Composite.Impl(right: IndexedTwo, rightLength) => - Buf.Composite.Impl( + case Buf.Composite(right: Vector[Buf], rightLength) => + Buf.Composite(left.a +: left.b +: left.c +: right, leftLength + rightLength) + case Buf.Composite(right: IndexedTwo, rightLength) => + Buf.Composite( (new VectorBuilder[Buf] += left.a += left.b += left.c += right.a += right.b) .result(), leftLength + rightLength ) - case Buf.Composite.Impl(right: IndexedThree, rightLength) => - Buf.Composite.Impl( - (new VectorBuilder[Buf] += left.a += left.b += left.c += right.a += right.b += right.c).result(), + case Buf.Composite(right: IndexedThree, rightLength) => + Buf.Composite( + (new VectorBuilder[ + Buf + ] += left.a += left.b += left.c += right.a += right.b += right.c) + .result(), leftLength + rightLength ) case _ => - Buf.Composite.Impl( + Buf.Composite( (new VectorBuilder[Buf] += left.a += left.b += left.c += right).result(), leftLength + right.length ) @@ -151,17 +159,17 @@ abstract class Buf { outer => case left => right match { - case Buf.Composite.Impl(right: Vector[Buf], rightLength) => - Buf.Composite.Impl(left +: right, left.length + rightLength) - case Buf.Composite.Impl(right: IndexedTwo, rightLength) => - Buf.Composite.Impl(new IndexedThree(left, right.a, right.b), left.length + rightLength) - case Buf.Composite.Impl(right: IndexedThree, rightLength) => - Buf.Composite.Impl( + case Buf.Composite(right: Vector[Buf], rightLength) => + Buf.Composite(left +: right, left.length + rightLength) + case Buf.Composite(right: IndexedTwo, rightLength) => + Buf.Composite(new IndexedThree(left, right.a, right.b), left.length + rightLength) + case Buf.Composite(right: IndexedThree, rightLength) => + Buf.Composite( (new VectorBuilder[Buf] += left += right.a += right.b += right.c).result(), left.length + rightLength ) case _ => - Buf.Composite.Impl(new IndexedTwo(left, right), left.length + right.length) + Buf.Composite(new IndexedTwo(left, right), left.length + right.length) } } } @@ -257,11 +265,8 @@ abstract class Buf { outer => protected[this] def isSliceIdentity(from: Int, until: Int): Boolean = from == 0 && until >= length - /** Helps implementations validate the arguments to [[write]]. */ - protected[this] def checkWriteArgs( - outputLen: Int, - outputOff: Int - ): Unit = { + /** Helps implementations validate the arguments to [[Buf.write(output:Array[Byte],off:Int):Unit*]]. */ + protected[this] def checkWriteArgs(outputLen: Int, outputOff: Int): Unit = { if (outputOff < 0) throw new IllegalArgumentException(s"offset must be non-negative: $outputOff") val len = length @@ -274,16 +279,20 @@ abstract class Buf { outer => } /** - * Buf wrapper-types (like Buf.ByteArray and Buf.ByteBuffer) provide Shared and - * Owned APIs, each of which with construction & extraction utilities. + * [[Buf]] wrapper-types (like [[Buf.ByteArray]] and [[Buf.ByteBuffer]]) provide + * "shared" and "owned" APIs, each of which contain construction and extraction + * utilities. While the owned and shared nomenclature is unintuitive, their + * purpose and use cases are straightforward. * - * The Owned APIs may provide direct access to a Buf's underlying - * implementation; and so mutating the data structure invalidates a Buf's - * immutability constraint. Users must take care to handle this data - * immutably. + * A "shared" [[Buf]] means that steps are taken to ensure that the produced + * [[Buf]] shares no state with the producer. This typically implies defensive + * copying and should be used when you do not control the lifecycle or usage + * of the passed in data. * - * The Shared variants, on the other hand, ensure that the Buf shares no state - * with the caller (at the cost of additional allocation). + * The "owned" APIs may provide direct access to a [[Buf]]'s underlying + * implementation; and so mutating the underlying data structure invalidates + * a [[Buf]]'s immutability constraint. Users must take care to handle this data + * immutably. * * Note: There are Java-friendly APIs for this object at `com.twitter.io.Bufs`. */ @@ -293,7 +302,7 @@ object Buf { * @return `true` if the processor would like to continue processing * more bytes and `false` otherwise. * - * @see [[Buf.process]] + * @see [[com.twitter.io.Buf.process(processor:com\.twitter\.io\.Buf\.Processor):Int* process(Processor)]] * * @note this is not a `Function1[Byte, Boolean]` despite very * much fitting that interface. This was done to avoiding boxing @@ -316,7 +325,10 @@ object Buf { checkSliceArgs(from, until) this } - protected def unsafeByteArrayBuf: Option[Buf.ByteArray] = None + protected val unsafeByteArrayBuf: Option[Buf.ByteArray] = Some( + new ByteArray(Array.emptyByteArray, 0, 0)) + override def unsafeByteArray: Array[Byte] = Array.emptyByteArray + def get(index: Int): Byte = throw new IndexOutOfBoundsException(s"Index out of bounds: $index") def process(from: Int, until: Int, processor: Processor): Int = { @@ -337,10 +349,10 @@ object Buf { val builder = Vector.newBuilder[Buf] var length = 0 bufs.foreach { - case b: Composite.Impl => + case Composite(b, l) => // Guaranteed to be non-empty by construction - length += b.length - builder ++= b.bufs + length += l + builder ++= b case b => val len = b.length @@ -356,7 +368,7 @@ object Buf { else if (filtered.size == 1) filtered.head else - new Composite.Impl(filtered, length) + Composite(filtered, length) } /** Helps Buf implementations validate the arguments to slicing functions. */ @@ -392,20 +404,10 @@ object Buf { def length: Int = 3 } - /** - * A `Buf` which is composed of other `Bufs`. - * - * @see [[Buf.apply]] for creating new instances. - */ - sealed abstract class Composite extends Buf { - def bufs: IndexedSeq[Buf] - protected def computedLength: Int - + final case class Composite(bufs: IndexedSeq[Buf], length: Int) extends Buf { // the factory method requires non-empty Bufs override def isEmpty: Boolean = false - def length: Int = computedLength - override def toString: String = s"Buf.Composite(length=$length)" def write(output: Array[Byte], off: Int): Unit = { @@ -432,7 +434,7 @@ object Buf { def slice(from: Int, until: Int): Buf = { checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) return Buf.Empty + if (this.isSliceEmpty(from, until)) return Buf.Empty else if (isSliceIdentity(from, until)) return this var begin = from @@ -488,6 +490,58 @@ object Buf { } } + def get(index: Int): Byte = { + var bufIdx = 0 + var byteIdx = 0 + while (bufIdx < bufs.length) { + val buf = bufs(bufIdx) + val bufLen = buf.length + if (index < byteIdx + bufLen) { + return buf.get(index - byteIdx) + } else { + byteIdx += bufLen + } + bufIdx += 1 + } + throw new IndexOutOfBoundsException(s"Index out of bounds: $index") + } + + def process(from: Int, until: Int, processor: Processor): Int = { + checkSliceArgs(from, until) + if (this.isSliceEmpty(from, until)) return -1 + var i = 0 + var bufIdx = 0 + var continue = true + while (continue && i < until && bufIdx < bufs.length) { + val buf = bufs(bufIdx) + val bufLen = buf.length + + if (i + bufLen < from) { + // skip ahead to the right Buf for `from` + bufIdx += 1 + i += bufLen + } else { + // ensure we are positioned correctly in the first Buf + var byteIdx = + if (i >= from) 0 + else from - i + val endAt = math.min(bufLen, until - i) + while (continue && byteIdx < endAt) { + val byte = buf.get(byteIdx) + if (processor(byte)) { + byteIdx += 1 + } else { + continue = false + } + } + bufIdx += 1 + i += byteIdx + } + } + if (continue) -1 + else i + } + protected def unsafeByteArrayBuf: Option[ByteArray] = None private[this] def equalsIndexed(other: Buf): Boolean = { @@ -511,7 +565,7 @@ object Buf { override def equals(other: Any): Boolean = other match { case otherBuf: Buf if length == otherBuf.length => otherBuf match { - case Composite(otherBufs) => + case Composite(otherBufs, _) => // this is 2 nested loops, with the outer loop tracking which // Buf's they are on. The inner loop compares individual bytes across // the Bufs "segments". @@ -553,94 +607,25 @@ object Buf { } } - object Composite { - def unapply(buf: Composite): Option[IndexedSeq[Buf]] = - Some(buf.bufs) - - /** - * Basic implementation of a [[Buf]] created from n-`Bufs`. - */ - private[Buf] final case class Impl(bufs: IndexedSeq[Buf], protected val computedLength: Int) - extends Buf.Composite { - - bufs match { - case _: Vector[Buf] => - case _: IndexedTwo => - case _: IndexedThree => - case _ => throw new IllegalArgumentException(s"Type not allowed: ${bufs.getClass.getName}") - } - - // ensure there is a need for a `Composite` - if (bufs.length <= 1) - throw new IllegalArgumentException(s"Must have 2 or more bufs: $bufs") - if (computedLength <= 0) - throw new IllegalArgumentException(s"Length must be positive: $computedLength") - - // the factory method requires non-empty Bufs - override def isEmpty: Boolean = false - - def get(index: Int): Byte = { - var bufIdx = 0 - var byteIdx = 0 - while (bufIdx < bufs.length) { - val buf = bufs(bufIdx) - val bufLen = buf.length - if (index < byteIdx + bufLen) { - return buf.get(index - byteIdx) - } else { - byteIdx += bufLen - } - bufIdx += 1 - } - throw new IndexOutOfBoundsException(s"Index out of bounds: $index") - } - - def process(from: Int, until: Int, processor: Processor): Int = { - checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) return -1 - var i = 0 - var bufIdx = 0 - var continue = true - while (continue && i < until && bufIdx < bufs.length) { - val buf = bufs(bufIdx) - val bufLen = buf.length - - if (i + bufLen < from) { - // skip ahead to the right Buf for `from` - bufIdx += 1 - i += bufLen - } else { - // ensure we are positioned correctly in the first Buf - var byteIdx = - if (i >= from) 0 - else from - i - val endAt = math.min(bufLen, until - i) - while (continue && byteIdx < endAt) { - val byte = buf.get(byteIdx) - if (processor(byte)) { - byteIdx += 1 - } else { - continue = false - } - } - bufIdx += 1 - i += byteIdx - } - } - if (continue) -1 - else i - } - } - } - /** * A buffer representing an array of bytes. */ class ByteArray( private[Buf] val bytes: Array[Byte], private[Buf] val begin: Int, - private[Buf] val end: Int - ) extends Buf { + private[Buf] val end: Int) + extends Buf { + + /** + * The backing byte[]. The array is only valid from [unsafeArrayOffset, unsafeArrayOffset + length), any data + * outside of those bounds is undefined. + */ + def unsafeRawByteArray: Array[Byte] = bytes + + /** + * The start offset of the data in the array returned by `unsafeRawByteArray`. + */ + def unsafeArrayOffset: Int = begin def get(index: Int): Byte = { val off = begin + index @@ -651,7 +636,7 @@ object Buf { def process(from: Int, until: Int, processor: Processor): Int = { checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) return -1 + if (this.isSliceEmpty(from, until)) return -1 var i = from var continue = true val endAt = math.min(until, length) @@ -678,7 +663,7 @@ object Buf { def slice(from: Int, until: Int): Buf = { checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) Buf.Empty + if (this.isSliceEmpty(from, until)) Buf.Empty else if (isSliceIdentity(from, until)) this else { val cap = math.min(until, length) @@ -737,12 +722,15 @@ object Buf { /** * Construct a buffer representing the given bytes. + * + * The contract is that these bytes will not be modified after this has + * been called. */ def apply(bytes: Byte*): Buf = Owned(bytes.toArray) /** - * Safely coerce a buffer to a Buf.ByteArray, potentially without copying its underlying - * data. + * Safely coerce a buffer to a [[Buf.ByteArray]], potentially without + * copying its underlying data. */ def coerce(buf: Buf): Buf.ByteArray = buf match { case buf: Buf.ByteArray => buf @@ -755,7 +743,16 @@ object Buf { } } - /** Owned non-copying constructors/extractors for Buf.ByteArray. */ + /** + * Non-copying constructors and extractors for [[Buf.ByteArray]]. + * + * The "owned" APIs may provide direct access to a [[Buf]]'s underlying + * implementation; and so mutating the underlying data structure invalidates + * a [[Buf]]'s immutability constraint. Users must take care to handle this data + * immutably. + * + * @see [[Shared]] for analogs that make defensive copies. + */ object Owned { /** @@ -773,7 +770,7 @@ object Buf { def apply(bytes: Array[Byte]): Buf = apply(bytes, 0, bytes.length) /** Extract the buffer's underlying offsets and array of bytes. */ - def unapply(buf: ByteArray): Option[(Array[Byte], Int, Int)] = + def unapply(buf: ByteArray): Some[(Array[Byte], Int, Int)] = Some((buf.bytes, buf.begin, buf.end)) /** @@ -781,17 +778,28 @@ object Buf { * * A copy may be performed if necessary. */ - def extract(buf: Buf): Array[Byte] = Buf.ByteArray.coerce(buf) match { - case Buf.ByteArray.Owned(bytes, 0, end) if end == bytes.length => - bytes - case Buf.ByteArray.Shared(bytes) => + def extract(buf: Buf): Array[Byte] = { + val byteArray = Buf.ByteArray.coerce(buf) + if (byteArray.begin == 0 && byteArray.end == byteArray.bytes.length) { + byteArray.bytes + } else { // If the unsafe version included offsets, we need to create a new array // containing only the relevant bytes. - bytes + byteArray.copiedByteArray + } } } - /** Safe copying constructors / extractors for Buf.ByteArray. */ + /** + * Copying constructors and extractors for [[Buf.ByteArray]]. + * + * A "shared" [[Buf]] means that steps are taken to ensure that the produced + * [[Buf]] shares no state with the producer. This typically implies defensive + * copying and should be used when you do not control the lifecycle or usage + * of the passed in data. + * + * @see [[Owned]] for analogs that do not make defensive copies. + */ object Shared { /** @@ -812,7 +820,7 @@ object Buf { def apply(bytes: Array[Byte]): Buf = apply(bytes, 0, bytes.length) /** Extract a copy of the buffer's underlying array of bytes. */ - def unapply(ba: ByteArray): Option[Array[Byte]] = Some(ba.copiedByteArray) + def unapply(ba: ByteArray): Some[Array[Byte]] = Some(ba.copiedByteArray) /** Get a copy of a a Buf's data as an array of bytes. */ def extract(buf: Buf): Array[Byte] = Buf.ByteArray.coerce(buf).copiedByteArray @@ -821,7 +829,7 @@ object Buf { /** * A buffer representing the remaining bytes in the - * given ByteBuffer. The given buffer will not be + * given `java.nio.ByteBuffer`. The given buffer will not be * affected. * * Modifications to the ByteBuffer's content will be @@ -836,7 +844,7 @@ object Buf { def process(from: Int, until: Int, processor: Processor): Int = { checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) return -1 + if (this.isSliceEmpty(from, until)) return -1 val pos = underlying.position() var i = from var continue = true @@ -866,19 +874,21 @@ object Buf { def slice(from: Int, until: Int): Buf = { checkSliceArgs(from, until) - if (isSliceEmpty(from, until)) Buf.Empty + if (this.isSliceEmpty(from, until)) Buf.Empty else if (isSliceIdentity(from, until)) this else { val dup = underlying.duplicate() - val limit = dup.position + math.min(until, length) - if (dup.limit > limit) dup.limit(limit) - dup.position(dup.position + from) + val limit = dup.position() + math.min(until, length) + if (dup.limit() > limit) dup.limit(limit) + dup.position(dup.position() + from) new ByteBuffer(dup) } } override def equals(other: Any): Boolean = other match { - case ByteBuffer(otherBB) => underlying.equals(otherBB) + // while normally this might use `ByteBuffer.unapply` this + // approach avoids the unnecessary `java.nio.ByteBuffer` copies. + case bb: ByteBuffer => underlying.equals(bb.underlying) case buf: Buf => Buf.equals(this, buf) case _ => false } @@ -886,7 +896,7 @@ object Buf { protected def unsafeByteArrayBuf: Option[Buf.ByteArray] = if (underlying.hasArray) { val array = underlying.array - val begin = underlying.arrayOffset + underlying.position + val begin = underlying.arrayOffset + underlying.position() val end = begin + underlying.remaining Some(new ByteArray(array, begin, end)) } else None @@ -894,11 +904,11 @@ object Buf { object ByteBuffer { - /** Extract a read-only view of the underlying [[java.nio.ByteBuffer]]. */ - def unapply(buf: ByteBuffer): Option[java.nio.ByteBuffer] = + /** Extract a read-only view of the underlying `java.nio.ByteBuffer`. */ + def unapply(buf: ByteBuffer): Some[java.nio.ByteBuffer] = Some(buf.underlying.asReadOnlyBuffer) - /** Coerce a generic buffer to a Buf.ByteBuffer, potentially without copying data. */ + /** Coerce a generic buffer to a [[Buf.ByteBuffer]], potentially without copying data. */ def coerce(buf: Buf): ByteBuffer = buf match { case buf: ByteBuffer => buf case _ => @@ -911,21 +921,30 @@ object Buf { new ByteBuffer(bb) } - /** Owned non-copying constructors/extractors for Buf.ByteBuffer. */ + /** + * Non-copying constructors and extractors for [[Buf.ByteBuffer]]. + * + * The "owned" APIs may provide direct access to a [[Buf]]'s underlying + * implementation; and so mutating the underlying data structure invalidates + * a [[Buf]]'s immutability constraint. Users must take care to handle this data + * immutably. + * + * @see [[Shared]] for analogs that make defensive copies. + */ object Owned { // N.B. We cannot use ByteBuffer.asReadOnly to ensure correctness because // it prevents direct access to its underlying byte array. /** - * Create a Buf.ByteBuffer by directly wrapping the provided [[java.nio.ByteBuffer]]. + * Create a Buf.ByteBuffer by directly wrapping the provided `java.nio.ByteBuffer`. */ def apply(bb: java.nio.ByteBuffer): Buf = if (bb.remaining == 0) Buf.Empty else new ByteBuffer(bb) - /** Extract the buffer's underlying [[java.nio.ByteBuffer]]. */ - def unapply(buf: ByteBuffer): Option[java.nio.ByteBuffer] = Some(buf.underlying) + /** Extract the buffer's underlying `java.nio.ByteBuffer`. */ + def unapply(buf: ByteBuffer): Some[java.nio.ByteBuffer] = Some(buf.underlying) /** * Get a reference to a Buf's data as a ByteBuffer. @@ -935,7 +954,16 @@ object Buf { def extract(buf: Buf): java.nio.ByteBuffer = Buf.ByteBuffer.coerce(buf).underlying } - /** Safe copying constructors/extractors for Buf.ByteBuffer. */ + /** + * Copying constructors and extractors for [[Buf.ByteBuffer]]. + * + * A "shared" [[Buf]] means that steps are taken to ensure that the produced + * [[Buf]] shares no state with the producer. This typically implies defensive + * copying and should be used when you do not control the lifecycle or usage + * of the passed in data. + * + * @see [[Owned]] for analogs that do not make defensive copies. + */ object Shared { private[this] def copy(orig: java.nio.ByteBuffer): java.nio.ByteBuffer = { val copy = java.nio.ByteBuffer.allocate(orig.remaining) @@ -987,7 +1015,7 @@ object Buf { def hash(buf: Buf): Int = finishHash(hashBuf(buf)) // Adapted from util-hashing. - private[this] val UintMax: Long = 0xFFFFFFFFL + private[this] val UintMax: Long = 0xffffffffL private[this] val Fnv1a32Prime: Int = 16777619 private[this] val Fnv1a32Init: Long = 0x811c9dc5L private[this] def finishHash(hash: Long): Int = (hash & UintMax).toInt @@ -1032,11 +1060,34 @@ object Buf { digits.toString } + /** + * Convert a hex string (eg one from slowHexString above) + * to a Buf of the string contents interpreted as a string of hexadecimal numbers. + * + * If an invalid string is passed to this method, the return value is not defined. It may throw + * or it may return a bogus Buf value. + */ + def slowFromHexString(hex: String): Buf = { + if (hex.length % 2 != 0) { + throw new IllegalArgumentException(s"Hex string is invalid: '$hex'") + } + val bytes = new Array[Byte](hex.length / 2) + var bytesIndex = 0 + while (bytesIndex < bytes.length) { + val hexIndex = bytesIndex * 2 + val b = Integer.parseInt(hex.substring(hexIndex, hexIndex + 2), 16) + bytes(bytesIndex) = b.toByte + bytesIndex = bytesIndex + 1 + } + + Buf.ByteArray.Owned(bytes) + } + /** * Create and deconstruct Utf-8 encoded buffers. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.UTF_8` for a Java-friendly API. */ @@ -1046,7 +1097,7 @@ object Buf { * Create and deconstruct 16-bit UTF buffers. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.UTF_16` for a Java-friendly API. */ @@ -1057,7 +1108,7 @@ object Buf { * with big-endian byte order. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.UTF_16BE` for a Java-friendly API. */ @@ -1068,7 +1119,7 @@ object Buf { * with little-endian byte order. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.UTF_16LE` for a Java-friendly API. */ @@ -1079,7 +1130,7 @@ object Buf { * ISO Latin Alphabet No. 1 charset. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.ISO_8859_1` for a Java-friendly API. */ @@ -1091,19 +1142,19 @@ object Buf { * Unicode character set. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @note See `com.twitter.io.Bufs.US_ASCII` for a Java-friendly API. */ object UsAscii extends StringCoder(JChar.US_ASCII) /** - * A [[StringCoder]] for a given [[java.nio.charset.Charset]] provides an - * [[apply(String) encoder]]: `String` to [[Buf]] - * and an [[unapply(Buf) extractor]]: [[Buf]] to `Option[String]`. + * A [[StringCoder]] for a given `java.nio.charset.Charset` provides an + * [[apply encoder]]: `String` to [[Buf]] + * and an [[unapply extractor]]: [[Buf]] to `Some[String]`. * * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` * * @see [[Utf8]] for UTF-8 encoding and decoding, and `Bufs.UTF8` for * Java users. Constants exist for other standard charsets as well. @@ -1122,10 +1173,9 @@ object Buf { * @note This extractor does *not* return None to indicate a failed * or impossible decoding. Malformed or unmappable bytes will * instead be silently replaced by the replacement character - * ("\uFFFD") in the returned String. This behavior may change - * in the future. + * ("\uFFFD") in the returned String. */ - def unapply(buf: Buf): Option[String] = Some(decodeString(buf, charset)) + def unapply(buf: Buf): Some[String] = Some(decodeString(buf, charset)) } /** @@ -1211,6 +1261,65 @@ object Buf { } } + /** + * Create and deconstruct unsigned 128-bit + * big endian encoded buffers. + * + * Deconstructing will return the value + * as well as the remaining buffer. + */ + object U128BE { + def apply(highBits: Long, lowBits: Long): Buf = { + val arr = new Array[Byte](16) + arr(0) = ((highBits >> 56) & 0xff).toByte + arr(1) = ((highBits >> 48) & 0xff).toByte + arr(2) = ((highBits >> 40) & 0xff).toByte + arr(3) = ((highBits >> 32) & 0xff).toByte + arr(4) = ((highBits >> 24) & 0xff).toByte + arr(5) = ((highBits >> 16) & 0xff).toByte + arr(6) = ((highBits >> 8) & 0xff).toByte + arr(7) = (highBits & 0xff).toByte + arr(8) = ((lowBits >> 56) & 0xff).toByte + arr(9) = ((lowBits >> 48) & 0xff).toByte + arr(10) = ((lowBits >> 40) & 0xff).toByte + arr(11) = ((lowBits >> 32) & 0xff).toByte + arr(12) = ((lowBits >> 24) & 0xff).toByte + arr(13) = ((lowBits >> 16) & 0xff).toByte + arr(14) = ((lowBits >> 8) & 0xff).toByte + arr(15) = (lowBits & 0xff).toByte + ByteArray.Owned(arr) + } + + // returns (highBits, lowBits, Buf) + def unapply(buf: Buf): Option[(Long, Long, Buf)] = + if (buf.length < 16) None + else { + val arr = new Array[Byte](16) + buf.slice(0, 16).write(arr, 0) + val rem = buf.slice(16, buf.length) + + val highBits = + ((arr(0) & 0xff).toLong << 56) | + ((arr(1) & 0xff).toLong << 48) | + ((arr(2) & 0xff).toLong << 40) | + ((arr(3) & 0xff).toLong << 32) | + ((arr(4) & 0xff).toLong << 24) | + ((arr(5) & 0xff).toLong << 16) | + ((arr(6) & 0xff).toLong << 8) | + (arr(7) & 0xff).toLong + val lowBits = + ((arr(8) & 0xff).toLong << 56) | + ((arr(9) & 0xff).toLong << 48) | + ((arr(10) & 0xff).toLong << 40) | + ((arr(11) & 0xff).toLong << 32) | + ((arr(12) & 0xff).toLong << 24) | + ((arr(13) & 0xff).toLong << 16) | + ((arr(14) & 0xff).toLong << 8) | + (arr(15) & 0xff).toLong + Some((highBits, lowBits, rem)) + } + } + /** * Create and deconstruct unsigned 32-bit * little endian encoded buffers. @@ -1285,6 +1394,65 @@ object Buf { } } + /** + * Create and deconstruct unsigned 128-bit + * little endian encoded buffers. + * + * Deconstructing will return the value + * as well as the remaining buffer. + */ + object U128LE { + def apply(highBits: Long, lowBits: Long): Buf = { + val arr = new Array[Byte](16) + arr(0) = (lowBits & 0xff).toByte + arr(1) = ((lowBits >> 8) & 0xff).toByte + arr(2) = ((lowBits >> 16) & 0xff).toByte + arr(3) = ((lowBits >> 24) & 0xff).toByte + arr(4) = ((lowBits >> 32) & 0xff).toByte + arr(5) = ((lowBits >> 40) & 0xff).toByte + arr(6) = ((lowBits >> 48) & 0xff).toByte + arr(7) = ((lowBits >> 56) & 0xff).toByte + arr(8) = (highBits & 0xff).toByte + arr(9) = ((highBits >> 8) & 0xff).toByte + arr(10) = ((highBits >> 16) & 0xff).toByte + arr(11) = ((highBits >> 24) & 0xff).toByte + arr(12) = ((highBits >> 32) & 0xff).toByte + arr(13) = ((highBits >> 40) & 0xff).toByte + arr(14) = ((highBits >> 48) & 0xff).toByte + arr(15) = ((highBits >> 56) & 0xff).toByte + ByteArray.Owned(arr) + } + + // returns (higherBits, lowerBits, Buf) + def unapply(buf: Buf): Option[(Long, Long, Buf)] = + if (buf.length < 16) None + else { + val arr = new Array[Byte](16) + buf.slice(0, 16).write(arr, 0) + val rem = buf.slice(16, buf.length) + + val lowBits = + (arr(0) & 0xff).toLong | + ((arr(1) & 0xff).toLong << 8) | + ((arr(2) & 0xff).toLong << 16) | + ((arr(3) & 0xff).toLong << 24) | + ((arr(4) & 0xff).toLong << 32) | + ((arr(5) & 0xff).toLong << 40) | + ((arr(6) & 0xff).toLong << 48) | + ((arr(7) & 0xff).toLong << 56) + val highBits = + (arr(8) & 0xff).toLong | + ((arr(9) & 0xff).toLong << 8) | + ((arr(10) & 0xff).toLong << 16) | + ((arr(11) & 0xff).toLong << 24) | + ((arr(12) & 0xff).toLong << 32) | + ((arr(13) & 0xff).toLong << 40) | + ((arr(14) & 0xff).toLong << 48) | + ((arr(15) & 0xff).toLong << 56) + Some((highBits, lowBits, rem)) + } + } + // Modeled after the Vector.++ operator, but use indexes instead of `for` // comprehensions for the optimized paths to save allocations and time. // See Scala ticket SI-7725 for details. diff --git a/util-core/src/main/scala/com/twitter/io/BufInputStream.scala b/util-core/src/main/scala/com/twitter/io/BufInputStream.scala index 25f3af3c7e..0a8941ae6f 100644 --- a/util-core/src/main/scala/com/twitter/io/BufInputStream.scala +++ b/util-core/src/main/scala/com/twitter/io/BufInputStream.scala @@ -25,7 +25,7 @@ class BufInputStream(val buf: Buf) extends InputStream { override def close(): Unit = {} // Marks the current position in this input stream. - override def mark(readlimit: Int) = synchronized { mrk = index } + override def mark(readlimit: Int): Unit = synchronized { mrk = index } // Tests if this input stream supports the mark and reset methods. override def markSupported(): Boolean = true @@ -36,7 +36,7 @@ class BufInputStream(val buf: Buf) extends InputStream { else { val b = buf.get(index) index += 1 - b & 0xFF + b & 0xff } } @@ -68,7 +68,7 @@ class BufInputStream(val buf: Buf) extends InputStream { * Repositions this stream to the position at the time the mark * method was last called on this input stream. */ - override def reset() = synchronized { index = mrk } + override def reset(): Unit = synchronized { index = mrk } /** * Skips over and discards n bytes of data from this input stream. diff --git a/util-core/src/main/scala/com/twitter/io/BufReader.scala b/util-core/src/main/scala/com/twitter/io/BufReader.scala index 9948eb39d9..d1a9415323 100644 --- a/util-core/src/main/scala/com/twitter/io/BufReader.scala +++ b/util-core/src/main/scala/com/twitter/io/BufReader.scala @@ -1,32 +1,124 @@ package com.twitter.io -import com.twitter.util.{Future, Return, Try, Throw} - -/** - * Construct a Reader from a Buf. - */ -private[io] class BufReader(buf: Buf) extends Reader[Buf] { - @volatile private[this] var state: Try[Buf] = Return(buf) - - def read(n: Int): Future[Option[Buf]] = synchronized { - state match { - case Return(buf) => - if (buf.isEmpty) Future.None +import com.twitter.util.Future +import java.util.NoSuchElementException +import scala.annotation.tailrec +import scala.collection.mutable.ListBuffer + +object BufReader { + + /** + * Creates a [[BufReader]] out of a given `buf`. The resulting [[Reader]] emits chunks of at + * most `Int.MaxValue` bytes. + */ + def apply(buf: Buf): Reader[Buf] = + apply(buf, Int.MaxValue) + + /** + * Creates a [[BufReader]] out of a given `buf`. The result [[Reader]] emits chunks of at most + * `chunkSize`. + */ + def apply(buf: Buf, chunkSize: Int): Reader[Buf] = { + // Returns an iterator that iterates over the given [[Buf]] by `chunkSize`. + def iterator(buf: Buf, chunkSize: Int): Iterator[Buf] = new Iterator[Buf] { + require(chunkSize > 0, s"Then chunkSize should be greater than 0, but received: $chunkSize") + private[this] var remainingBuf: Buf = buf + + def hasNext: Boolean = !remainingBuf.isEmpty + def next(): Buf = { + if (!hasNext) throw new NoSuchElementException("next() on empty iterator") + val nextChuck = remainingBuf.slice(0, chunkSize) + remainingBuf = remainingBuf.slice(chunkSize, Int.MaxValue) + nextChuck + } + } + + if (buf.isEmpty) Reader.empty[Buf] + else new IteratorReader[Buf](iterator(buf, chunkSize)) + } + + /** + * Read the entire bytestream presented by `r`. + */ + def readAll(r: Reader[Buf]): Future[Buf] = { + def loop(left: Buf): Future[Buf] = + r.read().flatMap { + case Some(right) => loop(left.concat(right)) + case _ => Future.value(left) + } + + loop(Buf.Empty) + } + + // see BufReader.chunked + private final class ChunkedFramer(chunkSize: Int) extends (Buf => Seq[Buf]) { + require(chunkSize > 0, s"chunkSize should be > 0 but was $chunkSize") + + def apply(input: Buf): Seq[Buf] = { + val acc = ListBuffer[Buf]() + + @tailrec + def loop(buf: Buf): Seq[Buf] = { + if (buf.length < chunkSize) (acc += buf).result() else { - val f = Future.value(Some(buf.slice(0, n))) - state = Return(buf.slice(n, Int.MaxValue)) - f + acc += buf.slice(0, chunkSize) + loop(buf.slice(chunkSize, buf.length)) } - case Throw(exc) => Future.exception(exc) + } + + loop(input) } } - def discard(): Unit = synchronized { - state = Throw(new Reader.ReaderDiscarded) + // see BufReader.framed + private final class Framed(r: Reader[Buf], framer: Buf => Seq[Buf]) + extends Reader[Buf] + with (Option[Buf] => Future[Option[Buf]]) { + + private[this] var frames: Seq[Buf] = Nil + + // we only enter here when `frames` is empty. + def apply(in: Option[Buf]): Future[Option[Buf]] = synchronized { + in match { + case Some(data) => + frames = framer(data) + read() + case None => + Future.None + } + } + + def read(): Future[Option[Buf]] = synchronized { + if (frames.isEmpty) { + // flatMap to `this` to prevent allocating + r.read().flatMap(this) + } else { + val nextFrame = frames.head + frames = frames.tail + Future.value(Some(nextFrame)) + } + } + + def discard(): Unit = synchronized { + frames = Seq.empty + r.discard() + } + + def onClose: Future[StreamTermination] = r.onClose } -} -object BufReader { - def apply(buf: Buf): Reader[Buf] = - if (buf.isEmpty) Reader.Null else new BufReader(buf) + /** + * Chunk the output of a given [[Reader]] by at most `chunkSize` (bytes). This consumes the + * reader. + */ + def chunked(r: Reader[Buf], chunkSize: Int): Reader[Buf] = + new Framed(r, new ChunkedFramer(chunkSize)) + + /** + * Wraps a [[Reader]] and emits frames as decided by `framer`. + * + * @note The returned `Reader` may not be thread safe depending on the behavior + * of the framer. + */ + def framed(r: Reader[Buf], framer: Buf => Seq[Buf]): Reader[Buf] = new Framed(r, framer) } diff --git a/util-core/src/main/scala/com/twitter/io/ByteReader.scala b/util-core/src/main/scala/com/twitter/io/ByteReader.scala index 0ffe75177a..7db3419546 100644 --- a/util-core/src/main/scala/com/twitter/io/ByteReader.scala +++ b/util-core/src/main/scala/com/twitter/io/ByteReader.scala @@ -184,14 +184,14 @@ trait ByteReader extends AutoCloseable { * Extract exactly the specified number of bytes into a `String` using the * specified `Charset`, advancing the byte cursor by `bytes`. * - * @throws UnderflowException if there are < `bytes` bytes available + * @throws ByteReader.UnderflowException if there are < `bytes` bytes available */ def readString(bytes: Int, charset: Charset): String /** * Skip over the next `n` bytes. * - * @throws UnderflowException if there are < `n` bytes available + * @throws ByteReader.UnderflowException if there are < `n` bytes available */ def skip(n: Int): Unit @@ -203,74 +203,6 @@ trait ByteReader extends AutoCloseable { def readAll(): Buf } -private[twitter] trait ProxyByteReader extends ByteReader { - protected def reader: ByteReader - - def remaining: Int = reader.remaining - - def remainingUntil(byte: Byte): Int = reader.remainingUntil(byte) - - def readByte(): Byte = reader.readByte() - - def readUnsignedByte(): Short = reader.readUnsignedByte() - - def readShortBE(): Short = reader.readShortBE() - - def readShortLE(): Short = reader.readShortLE() - - def readUnsignedShortBE(): Int = reader.readUnsignedShortBE() - - def readUnsignedShortLE(): Int = reader.readUnsignedShortLE() - - def readMediumBE(): Int = reader.readMediumBE() - - def readMediumLE(): Int = reader.readMediumLE() - - def readUnsignedMediumBE(): Int = reader.readUnsignedMediumBE() - - def readUnsignedMediumLE(): Int = reader.readUnsignedMediumLE() - - def readIntBE(): Int = reader.readIntBE() - - def readIntLE(): Int = reader.readIntLE() - - def readUnsignedIntBE(): Long = reader.readUnsignedIntBE() - - def readUnsignedIntLE(): Long = reader.readUnsignedIntLE() - - def readLongBE(): Long = reader.readLongBE() - - def readLongLE(): Long = reader.readLongLE() - - def readUnsignedLongBE(): BigInt = reader.readUnsignedLongBE() - - def readUnsignedLongLE(): BigInt = reader.readUnsignedLongLE() - - def readFloatBE(): Float = reader.readFloatBE() - - def readFloatLE(): Float = reader.readFloatLE() - - def readDoubleBE(): Double = reader.readDoubleBE() - - def readDoubleLE(): Double = reader.readDoubleLE() - - def readBytes(n: Int): Buf = reader.readBytes(n) - - def readString(bytes: Int, charset: Charset): String = reader.readString(bytes, charset) - - def skip(n: Int): Unit = reader.skip(n) - - def readAll(): Buf = reader.readAll() - - def process(from: Int, until: Int, processor: Buf.Processor): Int = - reader.process(from, until, processor) - - def process(processor: Buf.Processor): Int = - reader.process(processor) - - def close(): Unit = reader.close() -} - object ByteReader { /** @@ -290,7 +222,7 @@ object ByteReader { /** * The max value of a signed 24 bit "medium" integer */ - private[io] val SignedMediumMax = 0x800000 + private[io] val SignedMediumMax: Int = 0x800000 } private class ByteReaderImpl(buf: Buf) extends ByteReader { diff --git a/util-core/src/main/scala/com/twitter/io/CachingActivitySource.scala b/util-core/src/main/scala/com/twitter/io/CachingActivitySource.scala new file mode 100644 index 0000000000..a49a1abe6c --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/CachingActivitySource.scala @@ -0,0 +1,47 @@ +package com.twitter.io + +import com.twitter.util._ +import java.lang.ref.{ReferenceQueue, WeakReference} +import java.util.HashMap + +/** + * A convenient wrapper for caching the results returned by the + * underlying ActivitySource. + */ +class CachingActivitySource[T](underlying: ActivitySource[T]) extends ActivitySource[T] { + + private[this] val refq = new ReferenceQueue[Activity[T]] + private[this] val forward = new HashMap[String, WeakReference[Activity[T]]] + private[this] val reverse = new HashMap[WeakReference[Activity[T]], String] + + /** + * A caching proxy to the underlying ActivitySource. Vars are cached by + * name, and are tracked with WeakReferences. + */ + def get(name: String): Activity[T] = synchronized { + gc() + Option(forward.get(name)) flatMap { wr => Option(wr.get()) } match { + case Some(v) => v + case None => + val v = underlying.get(name) + val ref = new WeakReference(v, refq) + forward.put(name, ref) + reverse.put(ref, name) + v + } + } + + /** + * Remove garbage collected cache entries. + */ + def gc(): Unit = synchronized { + var ref = refq.poll() + while (ref != null) { + val key = reverse.remove(ref) + if (key != null) + forward.remove(key) + + ref = refq.poll() + } + } +} diff --git a/util-core/src/main/scala/com/twitter/io/Charsets.scala b/util-core/src/main/scala/com/twitter/io/Charsets.scala index 42a738b952..88960097de 100644 --- a/util-core/src/main/scala/com/twitter/io/Charsets.scala +++ b/util-core/src/main/scala/com/twitter/io/Charsets.scala @@ -4,7 +4,7 @@ import java.nio.charset.{Charset, CharsetDecoder, CharsetEncoder, CodingErrorAct import java.util /** - * A set of [[java.nio.charset.Charset]] utilities. + * A set of `java.nio.charset.Charset` utilities. */ object Charsets { @@ -17,10 +17,10 @@ object Charsets { } /** - * Returns a cached thread-local [[java.nio.charset.CharsetEncoder]] for + * Returns a cached thread-local `java.nio.charset.CharsetEncoder` for * the specified `charset`. * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` */ def encoder(charset: Charset): CharsetEncoder = { var e = encoders.get.get(charset) @@ -34,10 +34,10 @@ object Charsets { } /** - * Returns a cached thread-local [[java.nio.charset.CharsetDecoder]] for + * Returns a cached thread-local `java.nio.charset.CharsetDecoder` for * the specified `charset`. * @note Malformed and unmappable input is silently replaced - * see [[java.nio.charset.CodingErrorAction.REPLACE]] + * see `java.nio.charset.CodingErrorAction.REPLACE` */ def decoder(charset: Charset): CharsetDecoder = { var d = decoders.get.get(charset) diff --git a/util-core/src/main/scala/com/twitter/io/ClassLoaderActivitySource.scala b/util-core/src/main/scala/com/twitter/io/ClassLoaderActivitySource.scala new file mode 100644 index 0000000000..b29bf4e412 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/ClassLoaderActivitySource.scala @@ -0,0 +1,46 @@ +package com.twitter.io + +import com.twitter.util._ +import java.util.concurrent.atomic.AtomicBoolean + +/** + * An ActivitySource for ClassLoader resources. + */ +class ClassLoaderActivitySource private[io] (classLoader: ClassLoader, pool: FuturePool) + extends ActivitySource[Buf] { + + private[io] def this(classLoader: ClassLoader) = this(classLoader, FuturePool.unboundedPool) + + def get(name: String): Activity[Buf] = { + // This Var is updated at most once since ClassLoader + // resources don't change (do they?). + val runOnce = new AtomicBoolean(false) + val p = new Promise[Activity.State[Buf]] + + // Defer loading until the first observation + val v = Var.async[Activity.State[Buf]](Activity.Pending) { value => + if (runOnce.compareAndSet(false, true)) { + pool { + classLoader.getResourceAsStream(name) match { + case null => p.setValue(Activity.Failed(ActivitySource.NotFound)) + case stream => + val reader = InputStreamReader(stream, pool) + BufReader.readAll(reader) respond { + case Return(buf) => + p.setValue(Activity.Ok(buf)) + case Throw(cause) => + p.setValue(Activity.Failed(cause)) + } ensure { + // InputStreamReader ignores the deadline in close + reader.close(Time.Undefined) + } + } + } + } + p.onSuccess(value() = _) + Closable.nop + } + + Activity(v) + } +} diff --git a/util-core/src/main/scala/com/twitter/io/ClasspathResource.scala b/util-core/src/main/scala/com/twitter/io/ClasspathResource.scala new file mode 100644 index 0000000000..03f3d65f86 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/ClasspathResource.scala @@ -0,0 +1,46 @@ +package com.twitter.io + +import java.io.{BufferedInputStream, InputStream} + +object ClasspathResource { + + /** + * Loads the named resource from the classpath to return an optional [[InputStream]]. + * If the resolved resource exists and has a non-zero number of readable bytes an + * [[InputStream]] will be returned otherwise a None is returned. + * + * Resolution of the path to the named resource is ALWAYS assumed to be absolute, that is, + * if the given `name` does not begin with a `/` (`\u002f`), one will be prepended to the given name. + * E.g., `load("foo.txt")` will attempt to load the resource with the absolute name `/foo.txt`. + * + * ==Usage== + * {{{ + * load("foo.txt") // loads the classpath resource `/foo.txt` + * load("/foo.txt") // loads the classpath resource `/foo.txt` + * + * load("foo/bar/file.txt") // loads the classpath resource `/foo/bar/file.txt` + * load("/foo/bar/file.txt") // loads the classpath resource `/foo/bar/file.txt` + * }}} + * + * @note The caller is responsible for managing the closing of any returned [[InputStream]]. + * @see [[java.lang.Class#getResourceAsStream]] + */ + def load(name: String): Option[InputStream] = Option( + getClass.getResourceAsStream(absolutePath(name))) match { + case Some(inputStream) => + val bufferedInputStream: BufferedInputStream = new BufferedInputStream(inputStream) + if (bufferedInputStream.available() > 0) { + Some(bufferedInputStream) + } else { + // close the opened resource since we will not be returning the input stream + bufferedInputStream.close() + None + } + case _ => None + } + + // ensure the given name is prepended with a "/" to load as an absolute path. + private[this] def absolutePath(name: String): String = { + if (name.startsWith("/")) name else "/" + name + } +} diff --git a/util-core/src/main/scala/com/twitter/io/FilePollingActivitySource.scala b/util-core/src/main/scala/com/twitter/io/FilePollingActivitySource.scala new file mode 100644 index 0000000000..4ab0402f4a --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/FilePollingActivitySource.scala @@ -0,0 +1,47 @@ +package com.twitter.io + +import com.twitter.util._ +import java.io.{File, FileInputStream} + +/** + * An ActivitySource for observing the contents of a file with periodic polling. + */ +class FilePollingActivitySource private[io] ( + period: Duration, + pool: FuturePool +)( + implicit timer: Timer) + extends ActivitySource[Buf] { + + private[io] def this(period: Duration)(implicit timer: Timer) = + this(period, FuturePool.unboundedPool) + + def get(name: String): Activity[Buf] = { + val v = Var.async[Activity.State[Buf]](Activity.Pending) { value => + val timerTask = timer.schedule(Time.now, period) { + val file = new File(name) + + if (file.exists()) { + pool { + val reader = InputStreamReader(new FileInputStream(file), pool) + BufReader.readAll(reader) respond { + case Return(buf) => + value() = Activity.Ok(buf) + case Throw(cause) => + value() = Activity.Failed(cause) + } ensure { + // InputStreamReader ignores the deadline in close + reader.close(Time.Undefined) + } + } + } else { + value() = Activity.Failed(ActivitySource.NotFound) + } + } + + Closable.make { _ => Future { timerTask.cancel() } } + } + + Activity(v) + } +} diff --git a/util-core/src/main/scala/com/twitter/io/Files.scala b/util-core/src/main/scala/com/twitter/io/Files.scala index a2dcaf9e72..e8263eb317 100644 --- a/util-core/src/main/scala/com/twitter/io/Files.scala +++ b/util-core/src/main/scala/com/twitter/io/Files.scala @@ -42,9 +42,7 @@ object Files { } else if (file.isFile) { file.delete() } else { - file.listFiles.foreach { f => - delete(f) - } + file.listFiles.foreach { f => delete(f) } file.delete() } } diff --git a/util-core/src/main/scala/com/twitter/io/FutureReader.scala b/util-core/src/main/scala/com/twitter/io/FutureReader.scala new file mode 100644 index 0000000000..f7206c5a93 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/FutureReader.scala @@ -0,0 +1,104 @@ +package com.twitter.io + +import com.twitter.io.FutureReader._ +import com.twitter.util.Future +import com.twitter.util.Promise +import com.twitter.util.Return +import com.twitter.util.Throw + +/** + * We want to ensure that this reader always satisfies these invariants: + * 1. Satisfied closep is always aligned with the state + * 2. Reading from a discarded reader will always return ReaderDiscardedException + * 3. Reading from a fully read reader will always return None + * 4. Reading from a failed reader will always return the exception it has thrown + * + * We achieved this with a state machine where any access to the state is synchronized, + * and by preventing changing the state when the reader is failed, fully read, or discarded + */ +private[io] final class FutureReader[A](fa: Future[A]) extends Reader[A] { + private[this] val closep = Promise[StreamTermination]() + private[this] var state: State = State.Idle + + def read(): Future[Option[A]] = { + val (updatedState, result) = synchronized { + val result = state match { + case State.Idle => + fa.transform { + case Throw(t) => + state = State.Failed(t) + Future.exception(t) + case Return(v) => + state = State.Read + Future.value(Some(v)) + } + case State.Read => + state = State.FullyRead + Future.None + case State.Failed(_) => + // closep is guaranteed to be an exception, flatMap + // should never be triggered but return the exception + closep.flatMap(_ => Future.None) + case State.FullyRead => + Future.None + case State.Discarded => + Future.exception(new ReaderDiscardedException) + } + (state, result) + } + + updatedState match { + case State.Failed(t) => + // use `updateIfEmpty` instead of `update` here because we don't want to throw + // `ImmutableResult` exception when reading from a failed `FutureReader` multiple times + closep.updateIfEmpty(Throw(t)) + case State.FullyRead => + closep.updateIfEmpty(StreamTermination.FullyRead.Return) + case _ => + } + + result + } + + def discard(): Unit = { + val discarded = synchronized { + state match { + case State.Idle | State.Read => + state = State.Discarded + true + case _ => false + } + } + if (discarded) { + closep.updateIfEmpty(StreamTermination.Discarded.Return) + fa.raise(new ReaderDiscardedException) + } + } + + def onClose: Future[StreamTermination] = closep +} + +object FutureReader { + + /** + * Indicates reader state when the reader is created via FutureReader + */ + private sealed trait State + private object State { + + /** Indicates the reader is ready to be read. */ + case object Idle extends State + + /** Indicates the reader has been read. */ + case object Read extends State + + /** Indicates an exception occurred during reading */ + case class Failed(t: Throwable) extends State + + /** Indicates the EOS has been observed. */ + case object FullyRead extends State + + /** Indicates the reader has been discarded. */ + case object Discarded extends State + } +} diff --git a/util-core/src/main/scala/com/twitter/io/InputStreamReader.scala b/util-core/src/main/scala/com/twitter/io/InputStreamReader.scala index f47bf74b06..50e0028939 100644 --- a/util-core/src/main/scala/com/twitter/io/InputStreamReader.scala +++ b/util-core/src/main/scala/com/twitter/io/InputStreamReader.scala @@ -1,56 +1,63 @@ package com.twitter.io import com.twitter.concurrent.AsyncMutex -import com.twitter.util.{Closable, CloseAwaitably, Future, FuturePool, Time} +import com.twitter.util._ import java.io.InputStream /** * Provides the [[Reader]] API for an `InputStream`. * * The given `InputStream` will be closed when [[Reader.read]] - * reaches the EOF or a call to [[discard()]] or [[close()]]. + * reaches the EOF or a call to [[discard]] or `close`. */ -class InputStreamReader private[io] (inputStream: InputStream, maxBufferSize: Int, pool: FuturePool) +class InputStreamReader private[io] (inputStream: InputStream, chunkSize: Int, pool: FuturePool) extends Reader[Buf] with Closable with CloseAwaitably { private[this] val mutex = new AsyncMutex() @volatile private[this] var discarded = false + private[this] val closep = Promise[StreamTermination]() - def this(inputStream: InputStream, maxBufferSize: Int) = - this(inputStream, maxBufferSize, FuturePool.interruptibleUnboundedPool) + /** + * Constructs an [[InputStreamReader]] out of a given `inputStream`. The resulting [[Reader]] + * emits chunks of at most `chunkSize`. + */ + def this(inputStream: InputStream, chunkSize: Int) = + this(inputStream, chunkSize, FuturePool.interruptibleUnboundedPool) /** * Asynchronously read at most min(`n`, `maxBufferSize`) bytes from - * the `InputStream`. The returned [[Future]] represents the results of + * the `InputStream`. The returned [[com.twitter.util.Future]] represents the results of * the read operation. Any failure indicates an error; an empty buffer * indicates that the stream has completed. * * @note the underlying `InputStream` is closed on read of EOF. */ - def read(n: Int): Future[Option[Buf]] = { + def read(): Future[Option[Buf]] = { if (discarded) - return Future.exception(new Reader.ReaderDiscarded()) - if (n == 0) - return Future.value(Some(Buf.Empty)) + return Future.exception(new ReaderDiscardedException()) mutex.acquire().flatMap { permit => pool { try { if (discarded) - throw new Reader.ReaderDiscarded() - val size = math.min(n, maxBufferSize) - val buffer = new Array[Byte](size) - val c = inputStream.read(buffer, 0, size) + throw new ReaderDiscardedException() + val buffer = new Array[Byte](chunkSize) + val c = inputStream.read(buffer, 0, chunkSize) if (c == -1) { pool { inputStream.close() } + closep.updateIfEmpty(StreamTermination.FullyRead.Return) None } else { Some(Buf.ByteArray.Owned(buffer, 0, c)) } } catch { case exc: InterruptedException => - discard() + // we use updateIfEmpty because this is potentially racy, if someone + // called close and then we were interrupted. + if (closep.updateIfEmpty(Throw(exc))) { + discard() + } throw exc } }.ensure { @@ -62,18 +69,21 @@ class InputStreamReader private[io] (inputStream: InputStream, maxBufferSize: In /** * Discard this reader: its output is no longer required. * - * This closes the underlying `InputStream`. + * This asynchronously closes the underlying `InputStream`. */ - def discard(): Unit = - close() + def discard(): Unit = close() /** * Discards this [[Reader]] and closes the underlying `InputStream` */ def close(deadline: Time): Future[Unit] = closeAwaitably { discarded = true - pool { inputStream.close() } + pool { inputStream.close() }.ensure { + closep.updateIfEmpty(StreamTermination.Discarded.Return) + } } + + def onClose: Future[StreamTermination] = closep } object InputStreamReader { @@ -81,13 +91,18 @@ object InputStreamReader { /** * Create an [[InputStreamReader]] from the given `InputStream` - * using [[FuturePool.interruptibleUnboundedPool]] for executing - * all I/O. + * using [[com.twitter.util.FuturePool.interruptibleUnboundedPool]] + * for executing all I/O. + */ + def apply(inputStream: InputStream, chunkSize: Int = DefaultMaxBufferSize): InputStreamReader = + new InputStreamReader(inputStream, chunkSize) + + /** + * Create an [[InputStreamReader]] from the given `InputStream` + * using the provided [[com.twitter.util.FuturePool]] for executing all I/O. The resulting + * [[InputStreamReader]] emits chunks of at most `DefaultMaxBufferSize`. */ - def apply( - inputStream: InputStream, - maxBufferSize: Int = DefaultMaxBufferSize - ): InputStreamReader = - new InputStreamReader(inputStream, maxBufferSize) + def apply(inputStream: InputStream, pool: FuturePool): InputStreamReader = + new InputStreamReader(inputStream, chunkSize = DefaultMaxBufferSize, pool) } diff --git a/util-core/src/main/scala/com/twitter/io/IteratorReader.scala b/util-core/src/main/scala/com/twitter/io/IteratorReader.scala new file mode 100644 index 0000000000..8acc4e7bf9 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/IteratorReader.scala @@ -0,0 +1,76 @@ +package com.twitter.io + +import com.twitter.util.{Future, Promise} + +/** + * We want to ensure that this reader always satisfies these invariants: + * 1. Reading from a discarded reader will always return ReaderDiscardedException + * 2. Reading from a fully read reader will always return None + * + * We achieved this with a state machine where any access to the state is synchronized, + * and by preventing changing the state when the reader is fully read or discarded. + */ +private[io] final class IteratorReader[A](it: Iterator[A]) extends Reader[A] { + import IteratorReader._ + + private[this] val closep = Promise[StreamTermination]() + private[this] var state: State = State.Idle + + def read(): Future[Option[A]] = { + val result = synchronized { + state match { + case State.Idle => + if (it.hasNext) { + val n = it.next() + Future.value(Some(n)) + } else { + state = State.FullyRead + Future.None + } + case State.FullyRead => + Future.None + case State.Discarded => + Future.exception(new ReaderDiscardedException) + } + } + + if (result.eq(Future.None)) + closep.updateIfEmpty(StreamTermination.FullyRead.Return) + + result + } + + def discard(): Unit = { + val discarded = synchronized { + state match { + case State.Idle => + state = State.Discarded + true + case _ => false + } + } + if (discarded) closep.updateIfEmpty(StreamTermination.Discarded.Return) + + } + + def onClose: Future[StreamTermination] = closep +} + +private object IteratorReader { + + /** + * Indicates reader state when the reader is created via IteratorReader + */ + sealed trait State + object State { + + /** Indicates the reader is ready to be read. */ + case object Idle extends State + + /** Indicates the reader is fully read. */ + case object FullyRead extends State + + /** Indicates the reader has been discarded. */ + case object Discarded extends State + } +} diff --git a/util-core/src/main/scala/com/twitter/io/OutputStreamWriter.scala b/util-core/src/main/scala/com/twitter/io/OutputStreamWriter.scala index 23b2dbe8b1..f41400e5f2 100644 --- a/util-core/src/main/scala/com/twitter/io/OutputStreamWriter.scala +++ b/util-core/src/main/scala/com/twitter/io/OutputStreamWriter.scala @@ -1,7 +1,6 @@ package com.twitter.io -import com.twitter.io.Writer.ClosableWriter -import com.twitter.util.{Future, FuturePool, Promise, Return, Throw, Time} +import com.twitter.util.{Future, FuturePool, Promise, Return, Throw, Time, ConstFuture} import java.io.OutputStream import java.util.concurrent.atomic.AtomicReference import scala.annotation.tailrec @@ -9,10 +8,10 @@ import scala.annotation.tailrec /** * Construct a Writer from a given OutputStream. */ -private[io] class OutputStreamWriter(out: OutputStream, bufsize: Int) extends ClosableWriter[Buf] { +private[io] class OutputStreamWriter(out: OutputStream, bufsize: Int) extends Writer[Buf] { import com.twitter.io.OutputStreamWriter._ - private[this] val done = new Promise[Unit] + private[this] val done = new Promise[StreamTermination] private[this] val writeOp = new AtomicReference[Buf => Future[Unit]](doWrite) // Byte array reused on each write to avoid multiple allocations. @@ -37,8 +36,12 @@ private[io] class OutputStreamWriter(out: OutputStream, bufsize: Int) extends Cl buf => FuturePool.interruptibleUnboundedPool { drain(buf) } def write(buf: Buf): Future[Unit] = - if (done.isDefined) done - else + if (done.isDefined) { + done.transform { + case Return(StreamTermination.FullyRead) => Future.exception(CloseExc) + case t => (new ConstFuture(t)).unit + } + } else ( done or writeOp.getAndSet(_ => Future.exception(WriteExc))(buf) ) transform { @@ -55,13 +58,18 @@ private[io] class OutputStreamWriter(out: OutputStream, bufsize: Int) extends Cl def fail(cause: Throwable): Unit = done.updateIfEmpty(Throw(cause)) - def close(deadline: Time): Future[Unit] = - if (done.updateIfEmpty(Throw(CloseExc))) FuturePool.unboundedPool { + def close(deadline: Time): Future[Unit] = { + FuturePool.unboundedPool { out.close() - } else Future.Done + done.updateIfEmpty(StreamTermination.FullyRead.Return) + } + done.unit + } + + def onClose: Future[StreamTermination] = done } private object OutputStreamWriter { - val WriteExc = new IllegalStateException("write while writing") - val CloseExc = new IllegalStateException("write after close") + val WriteExc: IllegalStateException = new IllegalStateException("write while writing") + val CloseExc: IllegalStateException = new IllegalStateException("write after close") } diff --git a/util-core/src/main/scala/com/twitter/io/Pipe.scala b/util-core/src/main/scala/com/twitter/io/Pipe.scala index 5ebcdccc93..9a80c38490 100644 --- a/util-core/src/main/scala/com/twitter/io/Pipe.scala +++ b/util-core/src/main/scala/com/twitter/io/Pipe.scala @@ -1,151 +1,271 @@ package com.twitter.io -import com.twitter.util.{Closable, Future, Promise, Return, Time} +import com.twitter.util.{Future, Promise, Return, Throw, Time, Timer} /** * A synchronous in-memory pipe that connects [[Reader]] and [[Writer]] in the sense * that a reader's input is the output of a writer. * + * A pipe is structured as a smash of both interfaces, a [[Reader]] and a [[Writer]] such that can + * be passed directly to a consumer or a producer. + * + * {{{ + * def consumer(r: Reader[Buf]): Future[Unit] = ??? + * def producer(w: Writer[Buf]): Future[Unit] = ??? + * + * val p = new Pipe[Buf] + * + * consumer(p) + * producer(p) + * }}} + * * Reads and writes on the pipe are matched one to one and only one outstanding `read` - * or `write` is permitted in the current implementation. That is, the `write` (its - * returned [[Future]]) is resolved when the `read` consumes the written data. + * or `write` is permitted in the current implementation (multiple pending writes or reads + * resolve into `IllegalStateException` while leaving the pipe healthy). That is, the `write` + * (its returned [[com.twitter.util.Future]]) is resolved when the `read` consumes the written data. + * + * Here is, for example, a very typical write-loop that writes into a pipe-backed [[Writer]]: + * + * {{{ + * def writeLoop(w: Writer[Buf], data: List[Buf]): Future[Unit] = data match { + * case h :: t => p.write(h).before(writeLoop(w, t)) + * case Nil => w.close() + * } + * }}} + * + * Reading from a pipe-backed [[Reader]] is no different from working with any other reader: + * + *{{{ + * def readLoop(r: Reader[Buf], process: Buf => Future[Unit]): Future[Unit] = r.read().flatMap { + * case Some(chunk) => process(chunk).before(readLoop(r, process)) + * case None => Future.Done + * } + * }}} + * + * == Thread Safety == + * + * It is safe to call `read`, `write`, `fail`, `discard`, and `close` concurrently. The individual + * calls are synchronized on the given [[Pipe]]. + * + * == Closing or Failing Pipes == + * + * Besides expecting a write or a read, a pipe can be closed or failed. A writer can do both `close` + * and `fail` the pipe, while reader can only fail the pipe via `discard`. + * + * The following behavior is expected with regards to reading from or writing into a closed or a + * failed pipe: + * + * - Writing into a closed pipe yields `IllegalStateException` + * - Reading from a closed pipe yields EOF ([[com.twitter.util.Future.None]]) + * - Reading from or writing into a failed pipe yields a failure it was failed with + * + * It's also worth discussing how pipes are being closed. As closure is always initiated by a + * producer (writer), there is a machinery allowing it to be notified when said closure is observed + * by a consumer (reader). + * + * The following rules should help reasoning about closure signals in pipes: * - * It is safe to call `read`, `write`, and `close` concurrently. The individual calls - * are synchronized on the given [[Pipe]]. + * - Closing a pipe with a pending read resolves said read into EOF and returns a + * [[com.twitter.util.Future.Unit]] + * - Closing a pipe with a pending write by default fails said write with `IllegalStateException` and + * returns a future that will be satisfied when a consumer observes the closure (EOF) via read. If + * a timer is provided, the pipe will wait until the provided deadline for a successful read before + * failing the write. + * - Closing an idle pipe returns a future that will be satisfied when a consumer observes the + * closure (EOF) via read when a timer is provided, otherwise the pipe will be closed immedidately. */ -final class Pipe[A <: Buf] extends Reader[A] with Writer[A] with Closable { +final class Pipe[A](timer: Timer) extends Reader[A] with Writer[A] { + + // Implementation Notes: + // + // - The the only mutable state of this pipe is `state` and the access to it is synchronized + // on `this`. + // + // - We do not run promises under a lock (`synchronized`) as it could lead to a deadlock if + // there are interleaved queue operations in the waiter closure. + // + // - Although promises are run without a lock, synchronized `state` transitions guarantee + // a promise that needs to be run is now owned exclusively by a caller, hence no concurrent + // updates will be racing with it. import Pipe._ + /** + * For Java compatability + */ + def this() = this(Timer.Nil) + // thread-safety provided by synchronization on `this` private[this] var state: State[A] = State.Idle - def read(n: Int): Future[Option[A]] = synchronized { - state match { - case State.Failing(exc) => - Future.exception(exc) - - case State.Closing(reof) => - state = State.Eof - reof.setDone() - Future.None - - case State.Eof => - Future.None - - case State.Idle => - val p = new Promise[Option[A]] - state = State.Reading(n, p) - p - - case State.Writing(buf, p) if buf.length <= n => - // pending write can fully fit into this requested read - state = State.Idle - p.setDone() - Future.value(Some(buf)) - - case State.Writing(buf, p) => - // pending write is larger than the requested read - state = State.Writing(buf.slice(n, buf.length).asInstanceOf[A], p) - Future.value(Some(buf.slice(0, n).asInstanceOf[A])) - - case State.Reading(_, _) => - Future.exception(new IllegalStateException("read() while Reading")) + // satisfied when a `read` observes the EOF (after calling close()) + private[this] val closep: Promise[StreamTermination] = new Promise[StreamTermination]() + + def read(): Future[Option[A]] = { + val (waiter, result) = synchronized { + state match { + case State.Failed(exc) => + (null, Future.exception(exc)) + + case State.Closed => + (null, Future.None) + + case State.Idle => + val p = new Promise[Option[A]] + state = State.Reading(p) + (null, p) + + case State.Writing(buf, p) => + state = State.Idle + (p, Future.value(Some(buf))) + + case State.Reading(_) => + (null, Future.exception(new IllegalStateException("read() while read is pending"))) + + case State.Closing(buf, p) => + val result = buf match { + case None => + state = State.Closed + Future.None + case _ => + state = State.Closing(None, null) + Future.value(buf) + } + (p, result) + } } + + if (waiter != null) waiter.setDone() + if (result eq Future.None) closep.updateIfEmpty(StreamTermination.FullyRead.Return) + result } - def discard(): Unit = synchronized { - val cause = new Reader.ReaderDiscarded() - fail(cause) + def discard(): Unit = { + fail(new ReaderDiscardedException(), discard = true) } - def write(buf: A): Future[Unit] = synchronized { - state match { - case State.Failing(exc) => - Future.exception(exc) - - case State.Eof | State.Closing(_) => - Future.exception(new IllegalStateException("write after close")) - - case State.Idle => - val p = new Promise[Unit]() - state = State.Writing(buf, p) - p - - case State.Reading(n, p) if n < buf.length => - // pending reader doesn't have enough space for this write - val nextp = new Promise[Unit]() - state = State.Writing(buf.slice(n, buf.length).asInstanceOf[A], nextp) - p.setValue(Some(buf.slice(0, n).asInstanceOf[A])) - nextp - - case State.Reading(n, p) => - // pending reader has enough space for the full write - state = State.Idle - p.setValue(Some(buf)) - Future.Done - - case State.Writing(_, _) => - Future.exception(new IllegalStateException("write while Writing")) + def write(buf: A): Future[Unit] = { + val (waiter, value, result) = synchronized { + state match { + case State.Failed(exc) => + (null, null, Future.exception(exc)) + + case State.Closed | State.Closing(_, _) => + (null, null, Future.exception(new IllegalStateException("write() while closed"))) + + case State.Idle => + val p = new Promise[Unit] + state = State.Writing(buf, p) + (null, null, p) + + case State.Reading(p: Promise[Option[A]]) => + // pending reader has enough space for the full write + state = State.Idle + // The Scala 3 compiler differentiates between the class type, `A` and the the defined + // type of Promise in `Reading`, due to `State` being covariant in it's type parameter + // but invariant in `Reading`. Let's explictly cast this type as we know it will be of + // `Promise[Option[A]]`. See CSL-11210 or https://github.com/lampepfl/dotty/issues/13126 + (p.asInstanceOf[Promise[Option[A]]], Some(buf), Future.Done) + + case State.Writing(_, _) => + ( + null, + null, + Future.exception(new IllegalStateException("write() while write is pending")) + ) + } } + + // The waiter and the value are mutually inclusive so just checking against waiter is adequate. + if (waiter != null) waiter.setValue(value) + result } - def fail(cause: Throwable): Unit = synchronized { - // We fix the state before inspecting it so we enter in a - // good state if setting promises recurses. - val oldState = state - oldState match { - case State.Eof | State.Failing(_) => - // do not update state to failing - case State.Idle => - state = State.Failing(cause) - case State.Closing(reof) => - state = State.Failing(cause) - reof.setException(cause) - case State.Reading(_, p) => - state = State.Failing(cause) - p.setException(cause) - case State.Writing(_, p) => - state = State.Failing(cause) - p.setException(cause) + def fail(cause: Throwable): Unit = fail(cause, discard = false) + + private[this] def fail(cause: Throwable, discard: Boolean): Unit = { + val (closer, reader, writer) = synchronized { + state match { + case State.Closed | State.Failed(_) => + // do not update state to failing + (null, null, null) + case State.Idle => + state = State.Failed(cause) + (closep, null, null) + case State.Closing(_, p) => + state = State.Failed(cause) + (closep, null, p) + case State.Reading(p) => + state = State.Failed(cause) + (closep, p, null) + case State.Writing(_, p) => + state = State.Failed(cause) + (closep, null, p) + } + } + + if (reader != null) reader.setException(cause) + else if (writer != null) writer.setException(cause) + + if (closer != null) { + if (discard) closer.updateIfEmpty(StreamTermination.Discarded.Return) + else closer.updateIfEmpty(Throw(cause)) } } - def close(deadline: Time): Future[Unit] = synchronized { - state match { - case State.Failing(t) => - Future.exception(t) + private def closeLater(deadline: Time): Future[Unit] = { + timer.doLater(Time.now.until(deadline)) { + Pipe.this.synchronized { + state match { + case State.Failed(_) | State.Closed => + case _ => state = State.Closed + } + } + } + } - case State.Eof => - Future.Done + def close(deadline: Time): Future[Unit] = { + val (reader, writer) = synchronized { + state match { + case State.Failed(_) | State.Closed | State.Closing(_, _) => + (null, null) - case State.Closing(p) => - p + case State.Idle => + state = State.Closing(None, null) + (null, null) - case State.Idle => - val reof = new Promise[Unit]() - state = State.Closing(reof) - reof + case State.Reading(p) => + state = State.Closed + (p, null) - case State.Reading(_, p) => - state = State.Eof - p.update(Return.None) - Future.Done + case State.Writing(buf, p) => + state = State.Closing(Some(buf), p) + (null, p) + } + } - case State.Writing(_, p) => - val reof = new Promise[Unit]() - state = State.Closing(reof) - p.setException(new IllegalStateException("close while write is pending")) - reof + if (reader != null) { + reader.update(Return.None) + closep.update(StreamTermination.FullyRead.Return) + } else if (writer != null) { + closeLater(deadline).respond { _ => + val exn = new IllegalStateException("close() while write is pending") + writer.updateIfEmpty(Throw(exn)) + closep.updateIfEmpty(Throw(exn)) + } } + + onClose.unit } + def onClose: Future[StreamTermination] = closep + override def toString: String = synchronized(s"Pipe(state=$state)") } object Pipe { - private sealed trait State[+A <: Buf] + private sealed trait State[+A] private object State { @@ -155,30 +275,59 @@ object Pipe { /** * Indicates a read is pending and is awaiting a `write`. * - * @param n number of bytes to read. * @param p when satisfied it indicates that this read has completed. */ - final case class Reading[A <: Buf](n: Int, p: Promise[Option[A]]) extends State[A] + final case class Reading[A](p: Promise[Option[A]]) extends State[A] /** - * Indicates a write of `buf` is pending to be `read`. + * Indicates a write of `value` is pending to be `read`. * - * @param buf the [[Buf]] to write. + * @param value the value to write. * @param p when satisfied it indicates that this write has been fully read. */ - final case class Writing[A <: Buf](buf: A, p: Promise[Unit]) extends State[A] + final case class Writing[A](value: A, p: Promise[Unit]) extends State[A] - /** Indicates this was `fail`-ed or `discard`-ed. */ - final case class Failing(exc: Throwable) extends State[Nothing] + /** Indicates the pipe was failed. */ + final case class Failed(exc: Throwable) extends State[Nothing] /** - * Indicates a close occurred while `Idle` — no reads or writes were pending. + * Indicates a close occurred while a write was pending. * - * @param reof satisfied when a `read` sees the EOF. - */ - final case class Closing(reof: Promise[Unit]) extends State[Nothing] + * @param value the pending write + * @param p when satisfied it indicates that this write has been fully read or the + * close deadline has expired before it has been fully read. + * */ + final case class Closing[A](value: Option[A], p: Promise[Unit]) extends State[A] /** Indicates the reader has seen the EOF. No more reads or writes are allowed. */ - case object Eof extends State[Nothing] + case object Closed extends State[Nothing] + } + + /** + * Copy elements from a Reader to a Writer. The Reader will be discarded if + * `copy` is cancelled (discarding the writer). The Writer is unmanaged, the caller + * is responsible for finalization and error handling, e.g.: + * + * {{{ + * Pipe.copy(r, w, n) ensure w.close() + * }}} + */ + def copy[A](r: Reader[A], w: Writer[A]): Future[Unit] = { + def loop(): Future[Unit] = + r.read().flatMap { + case None => Future.Done + case Some(elem) => w.write(elem) before loop() + } + + w.onClose.respond { + case Return(StreamTermination.Discarded) => r.discard() + case _ => () + } + val p = new Promise[Unit] + // We have to do this because discarding the writer doesn't interrupt read + // operations, it only fails the next write operation. + loop().proxyTo(p) + p.setInterruptHandler { case _ => r.discard() } + p } } diff --git a/util-core/src/main/scala/com/twitter/io/ProxyByteReader.scala b/util-core/src/main/scala/com/twitter/io/ProxyByteReader.scala new file mode 100644 index 0000000000..a07b001b60 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/ProxyByteReader.scala @@ -0,0 +1,74 @@ +package com.twitter.io + +import java.nio.charset.Charset + +/** + * A proxy [[ByteReader]] that forwards all calls to another [[ByteReader]]. + * This is useful if you want to wrap-but-modify an existing [[ByteReader]]. + */ +abstract class ProxyByteReader(underlying: ByteReader) extends ByteReader { + + def remaining: Int = underlying.remaining + + def remainingUntil(byte: Byte): Int = underlying.remainingUntil(byte) + + def readByte(): Byte = underlying.readByte() + + def readUnsignedByte(): Short = underlying.readUnsignedByte() + + def readShortBE(): Short = underlying.readShortBE() + + def readShortLE(): Short = underlying.readShortLE() + + def readUnsignedShortBE(): Int = underlying.readUnsignedShortBE() + + def readUnsignedShortLE(): Int = underlying.readUnsignedShortLE() + + def readMediumBE(): Int = underlying.readMediumBE() + + def readMediumLE(): Int = underlying.readMediumLE() + + def readUnsignedMediumBE(): Int = underlying.readUnsignedMediumBE() + + def readUnsignedMediumLE(): Int = underlying.readUnsignedMediumLE() + + def readIntBE(): Int = underlying.readIntBE() + + def readIntLE(): Int = underlying.readIntLE() + + def readUnsignedIntBE(): Long = underlying.readUnsignedIntBE() + + def readUnsignedIntLE(): Long = underlying.readUnsignedIntLE() + + def readLongBE(): Long = underlying.readLongBE() + + def readLongLE(): Long = underlying.readLongLE() + + def readUnsignedLongBE(): BigInt = underlying.readUnsignedLongBE() + + def readUnsignedLongLE(): BigInt = underlying.readUnsignedLongLE() + + def readFloatBE(): Float = underlying.readFloatBE() + + def readFloatLE(): Float = underlying.readFloatLE() + + def readDoubleBE(): Double = underlying.readDoubleBE() + + def readDoubleLE(): Double = underlying.readDoubleLE() + + def readBytes(n: Int): Buf = underlying.readBytes(n) + + def readString(bytes: Int, charset: Charset): String = underlying.readString(bytes, charset) + + def skip(n: Int): Unit = underlying.skip(n) + + def readAll(): Buf = underlying.readAll() + + def process(from: Int, until: Int, processor: Buf.Processor): Int = + underlying.process(from, until, processor) + + def process(processor: Buf.Processor): Int = + underlying.process(processor) + + def close(): Unit = underlying.close() +} diff --git a/util-core/src/main/scala/com/twitter/io/ProxyByteWriter.scala b/util-core/src/main/scala/com/twitter/io/ProxyByteWriter.scala index 2ea6d669a4..8d5a713d28 100644 --- a/util-core/src/main/scala/com/twitter/io/ProxyByteWriter.scala +++ b/util-core/src/main/scala/com/twitter/io/ProxyByteWriter.scala @@ -3,10 +3,10 @@ package com.twitter.io import java.nio.charset.Charset /** - * A simple proxy ByteWriter that forwards all calls to another ByteWriter. - * This is useful if you want to wrap-but-modify an existing ByteWriter. + * A proxy [[ByteWriter]] that forwards all calls to another [[ByteWriter]]. + * This is useful if you want to wrap-but-modify an existing [[ByteWriter]]. */ -private[twitter] abstract class ProxyByteWriter(underlying: ByteWriter) extends AbstractByteWriter { +abstract class ProxyByteWriter(underlying: ByteWriter) extends AbstractByteWriter { def writeString(string: CharSequence, charset: Charset): ProxyByteWriter.this.type = { underlying.writeString(string, charset) diff --git a/util-core/src/main/scala/com/twitter/io/Reader.scala b/util-core/src/main/scala/com/twitter/io/Reader.scala index 4050610f3e..c07bf8a0a7 100644 --- a/util-core/src/main/scala/com/twitter/io/Reader.scala +++ b/util-core/src/main/scala/com/twitter/io/Reader.scala @@ -1,204 +1,330 @@ package com.twitter.io import com.twitter.concurrent.AsyncStream +import com.twitter.util.Promise.Detachable import com.twitter.util._ -import java.io.{File, FileInputStream, FileNotFoundException, InputStream} +import java.io.File +import java.io.FileInputStream +import java.io.InputStream +import java.util.{Collection => JCollection} +import scala.collection.mutable +import scala.jdk.CollectionConverters._ /** - * A Reader represents a stream of `A`s. + * A reader exposes a pull-based API to model a potentially infinite stream of arbitrary elements. * - * Readers permit at most one outstanding read. + * Given the pull-based API, the consumer is responsible for driving the computation. A very + * typical code pattern to consume a reader is to use a read-loop: * - * @tparam A the type of objects produced by this reader + * {{{ + * def consume[A](r: Reader[A])(process: A => Future[Unit]): Future[Unit] = + * r.read().flatMap { + * case Some(a) => process(a).before(consume(r)(process)) + * case None => Future.Done // reached the end of the stream; no need to discard + * } + * }}} + * + * One way to reason about the read-loop idiom is to view it as a subscription to a publisher + * (reader) in the Pub/Sub terminology. Perhaps the major difference between readers and traditional + * publishers, is readers only allow one subscriber (read-loop). It's generally safer to assume the + * reader is fully consumed (stream is exhausted) once its read-loop is run. + * + * == Error Handling == + * + * Given the read-loop above, its returned [[com.twitter.util.Future]] could be used to observe both + * successful and unsuccessful outcomes. + * + * {{{ + * consume(stream)(processor).respond { + * case Return(()) => println("Consumed an entire stream successfully.") + * case Throw(e) => println(s"Encountered an error while consuming a stream: \$e") + * } + * }}} + * + * @note Once failed, a stream can not be restarted such that all future reads will resolve into a + * failure. There is no need to discard an already failed stream. + * + * == Resource Safety == + * + * One of the important implications of readers, and streams in general, is that they are prone to + * resource leaks unless fully consumed or discarded (failed). Specifically, readers backed by + * network connections MUST be discarded unless already consumed (EOF observed) to prevent + * connection leaks. + * + * @note The read-loop above, for example, exhausts the stream (observes EOF) hence does not have to + * discard it (stream). + * + * == Back Pressure == + * + * The pattern above leverages [[com.twitter.util.Future]] recursion to exert back-pressure via + * allowing only one outstanding read. It's usually a good idea to structure consumers this way + * (i.e., the next read isn't issued until the previous read finishes). This would always ensure a + * finer grained back-pressure in network systems allowing the consumers to artificially slow down + * the producers and not rely solely on transport and IO buffering. + * + * @note Whether or not multiple outstanding reads are allowed on a `Reader` type is an undefined + * behaviour but could be changed in a refinement. + * + * == Cancellations == + * + * If a consumer is no longer interested in the stream, it can discard it. Note a discarded reader + * (or stream) can not be restarted. + * + * {{{ + * def consumeN[A](r: Reader[A], n: Int)(process: A => Future[Unit]): Future[Unit] = + * if (n == 0) Future(r.discard()) + * else r.read().flatMap { + * case Some(a) => process(a).before(consumeN(r, n - 1)(process)) + * case None => Future.Done // reached the end of the stream; no need to discard + * } + * }}} */ -trait Reader[+A] { +trait Reader[+A] { self => /** - * Asynchronously read at most `n` bytes from the byte stream. The - * returned future represents the results of the read request. If - * the read fails, the Reader is considered failed -- future reads - * will also fail. + * Asynchronously read the next element of this stream. Returned [[com.twitter.util.Future]] + * will resolve into `Some(e)` when the element is available or into `None` when stream + * is exhausted. * - * A result of None indicates EOF. + * Stream failures are terminal such that all subsequent reads will resolve in failed + * [[com.twitter.util.Future]]s. */ - def read(n: Int): Future[Option[A]] + def read(): Future[Option[A]] /** - * Discard this reader: its output is no longer required. + * Discard this stream as its output is no longer required. This could be used to signal the + * producer of this stream similarly how [[com.twitter.util.Future.raise]] used to propagate interrupts across + * future chains. + * + * @note Although unnecessary, it's always safe to discard a fully-consumed stream. */ def discard(): Unit -} -object Reader { + /** + * A [[com.twitter.util.Future]] that resolves once this reader is closed upon reading + * of end-of-stream. + * + * If the result is a failed future, this indicates that it was not closed either by reading + * until the end of the stream nor by discarding. This is useful for any extra resource cleanup + * that you must do once the stream is no longer being used. + */ + def onClose: Future[StreamTermination] + + /** + * Construct a new Reader by applying `f` to every item read from this Reader + * @param f the function constructs a new Reader[B] from the value of this Reader.read + * + * @note All operations of the new Reader will be in sync with self Reader. Discarding one Reader + * will discard the other Reader. When one Reader's onClose resolves, the other Reader's + * onClose will be resolved immediately with the same value. + */ + final def flatMap[B](f: A => Reader[B]): Reader[B] = Reader.flatten(map(f)) - val Null: Reader[Nothing] = new Reader[Nothing] { - def read(n: Int): Future[Option[Nothing]] = Future.None - def discard(): Unit = () + /** + * Construct a new Reader by applying `f` to every item read from this Reader + * @param f the function transforms data of type A to B + * + * @note All operations of the new Reader will be in sync with self Reader. Discarding one Reader + * will discard the other Reader. When one Reader's onClose resolves, the other Reader's + * onClose will be resolved immediately with the same value. + */ + final def map[B](f: A => B): Reader[B] = new Reader[B] { + def read(): Future[Option[B]] = self.read().map(oa => oa.map(f)) + def discard(): Unit = self.discard() + def onClose: Future[StreamTermination] = self.onClose } - def empty[A]: Reader[A] = Null.asInstanceOf[Reader[A]] + /** + * Converts a `Reader[Reader[B]]` into a `Reader[B]` + * + * @note All operations of the new Reader will be in sync with the outermost Reader. Discarding + * one Reader will discard the other Reader. When one Reader's onClose resolves, the other + * Reader's onClose will be resolved immediately with the same value. + * The subsequent readers are unmanaged, the caller is responsible for discarding those + * when abandoned. + */ + def flatten[B](implicit ev: A <:< Reader[B]): Reader[B] = + Reader.flatten(this.asInstanceOf[Reader[Reader[B]]]) +} - // See Reader.chunked - private final class Chunked(r: Reader[Buf], chunkSize: Int) - extends Reader[Buf] with (Option[Buf] => Future[Option[Buf]]) { +/** + * Abstract `Reader` class for Java compatibility. + */ +abstract class AbstractReader[+A] extends Reader[A] - require(chunkSize > 0, s"chunkSize should be > 0 but was $chunkSize") +/** + * Indicates that a given stream was discarded by the Reader's consumer. + */ +class ReaderDiscardedException extends Exception("Reader's consumer has discarded the stream") - private[this] var state = Buf.Empty +object Reader { - // We only enter here when state (buffer) is empty. - def apply(in: Option[Buf]): Future[Option[Buf]] = synchronized { - in match { - case Some(b) => - state = b - read(chunkSize) - case None => Future.None - } + def empty[A]: Reader[A] = new Reader[A] { + private[this] val closep = Promise[StreamTermination]() + def read(): Future[Option[Nothing]] = { + if (closep.updateIfEmpty(StreamTermination.FullyRead.Return)) + Future.None + else + closep.flatMap { + case StreamTermination.FullyRead => Future.None + case StreamTermination.Discarded => Future.exception(new ReaderDiscardedException) + } } - - def read(n: Int): Future[Option[Buf]] = synchronized { - // flatMap to `this` to prevent allocating - if (state.isEmpty) r.read(Int.MaxValue).flatMap(this) - else { - val result = state.slice(0, chunkSize) - state = state.slice(chunkSize, Int.MaxValue) - Future.value(Some(result)) - } + def discard(): Unit = { + closep.updateIfEmpty(StreamTermination.Discarded.Return) + () } - - def discard(): Unit = r.discard() + def onClose: Future[StreamTermination] = closep } - // see Reader.framed - private final class Framed(r: Reader[Buf], framer: Buf => Seq[Buf]) - extends Reader[Buf] with (Option[Buf] => Future[Option[Buf]]) { + /** + * Construct a `Reader` from a `Future` + * + * @note Multiple outstanding reads are not allowed on this reader + */ + def fromFuture[A](fa: Future[A]): Reader[A] = new FutureReader(fa) - private[this] var frames: Seq[Buf] = Seq.empty + /** + * Construct a `Reader` from a value `a` + * @note Multiple outstanding reads are not allowed on this reader + */ + def value[A](a: A): Reader[A] = fromFuture[A](Future.value(a)) - // we only enter here when `frames` is empty. - def apply(in: Option[Buf]): Future[Option[Buf]] = synchronized { - in match { - case Some(data) => - frames = framer(data) - read(Int.MaxValue) - case None => - Future.None - } - } + /** + * Construct a `Reader` from an exception `e` + */ + def exception[A](e: Throwable): Reader[A] = fromFuture[A](Future.exception(e)) - def read(n: Int): Future[Option[Buf]] = synchronized { - frames match { - case nextFrame :: rst => - frames = rst - Future.value(Some(nextFrame)) - case _ => - // flatMap to `this` to prevent allocating - r.read(Int.MaxValue).flatMap(this) - } - } + private def failOrDiscard(r: Reader[_], t: Throwable): Unit = r match { + case p: Pipe[_] => p.fail(t) + case r => r.discard() + } - def discard(): Unit = synchronized { - frames = Seq.empty - r.discard() + private def consume[A](r: Reader[A], builder: mutable.Builder[A, IndexedSeq[A]]): Future[Seq[A]] = + r.read().flatMap { + case Some(v) => + builder += v + consume(r, builder) + case None => + Future.value(builder.result()) } - } /** - * Read the entire bytestream presented by `r`. + * Read all items from the Reader r. + * @param r The reader to read from + * @return A Sequence of items. */ - def readAll(r: Reader[Buf]): Future[Buf] = { - def loop(left: Buf): Future[Buf] = - r.read(Int.MaxValue).flatMap { - case Some(right) => loop(left concat right) - case _ => Future.value(left) - } + def readAllItems[A](r: Reader[A]): Future[Seq[A]] = { + // This could be implemented as `readAllItemsInterruptable(r).masked`, but doing so would cause multiple extra + // promise allocations. + consume(r, IndexedSeq.newBuilder) + } - loop(Buf.Empty) + /** + * Read all items from the Reader r, interrupting the returned future will discard the reader. + * @param r The reader to read from + * @return A Sequence of items. + */ + def readAllItemsInterruptible[A](r: Reader[A]): Future[Seq[A]] = { + val p = new Promise[Seq[A]] with Promise.InterruptHandler { + override protected def onInterrupt(t: Throwable): Unit = failOrDiscard(r, t) + } + consume(r, IndexedSeq.newBuilder).proxyTo(p) + p } /** - * Chunk the output of a given [[Reader]] by at most `chunkSize` (bytes). This consumes the - * reader. - * - * @note The `n` (number of bytes to read) argument on the returned reader is ignored - * (`Int.MaxValue` is used instead). + * Create a new [[Reader]] from a given [[Buf]]. The output of a returned reader is chunked by + * a least `chunkSize` (bytes). */ - def chunked(r: Reader[Buf], chunkSize: Int): Reader[Buf] = new Chunked(r, chunkSize) + def fromBuf(buf: Buf): Reader[Buf] = fromBuf(buf, Int.MaxValue) /** - * Reader from a Buf. + * Create a new [[Reader]] from a given [[Buf]]. The output of a returned reader is chunked by + * at most `chunkSize` (bytes). + * + * @note The `n` (number of bytes to read) argument on the returned reader's `read` is ignored. */ - def fromBuf(buf: Buf): Reader[Buf] = BufReader(buf) + def fromBuf(buf: Buf, chunkSize: Int): Reader[Buf] = BufReader(buf, chunkSize) - class ReaderDiscarded extends Exception("This writer's reader has been discarded") + /** + * Create a new [[Reader]] from a given `File`. The output of a returned reader is chunked by + * at most `chunkSize` (bytes). + * + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. + * + * @see `Readers.fromFile` for a Java API + */ + def fromFile(f: File): Reader[Buf] = fromFile(f, InputStreamReader.DefaultMaxBufferSize) /** - * A [[Reader]] that is linked with a [[Writer]] and `close`-ing - * is synchronous. + * Create a new [[Reader]] from a given `File`. The output of a returned reader is chunked by + * at most `chunkSize` (bytes). + * + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. * - * Just as with [[Reader readers]] and [[Writer writers]], - * only one outstanding `read` or `write` is permitted. + * @see `Readers.fromFile` for a Java API + */ + def fromFile(f: File, chunkSize: Int): Reader[Buf] = fromStream(new FileInputStream(f), chunkSize) + + /** + * Create a new [[Reader]] from a given `Iterator`. * - * For a proper `close`, it should only be done when - * no writes are outstanding: - * {{{ - * val rw = Reader.writable() - * ... - * rw.write(buf).before(rw.close()) - * }}} + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. * - * If a producer is interested in knowing when all writes - * have been read and the reader has seen the EOF, it can - * wait until the future returned by `close()` is satisfied: - * {{{ - * val rw = Reader.writable() - * ... - * rw.close().ensure { - * println("party on! ♪┏(・o・)┛♪ the Reader has seen the EOF") - * } - * }}} + * @note It is not recommended to call `it.next()` after creating a `Reader` from it. + * Doing so will affect the behavior of `Reader.read()` because it will skip + * the value returned from `it.next`. */ - @deprecated("Use Pipe[A] instead", "2018-8-7") - type Writable[A <: Buf] = Pipe[A] + def fromIterator[A](it: Iterator[A]): Reader[A] = new IteratorReader[A](it) /** - * Create a new [[Writable]] which is a [[Reader]] that is linked - * with a [[Writer]]. + * Create a new [[Reader]] from a given `InputStream`. The output of a returned reader is + * chunked by at most `chunkSize` (bytes). * - * @see Readers.writable() for a Java API. + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. */ - @deprecated("Use Pipe() instead", "2018-8-7") - def writable(): Pipe[Buf] = new Pipe[Buf] + def fromStream(s: InputStream): Reader[Buf] = + fromStream(s, InputStreamReader.DefaultMaxBufferSize) /** - * Create a new [[Reader]] for a `File`. + * Create a new [[Reader]] from a given `InputStream`. The output of a returned reader is + * chunked by at most `chunkSize` (bytes). * - * The resources held by the returned [[Reader]] are released - * on reading of EOF and [[Reader.discard()]]. + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. * - * @see `Readers.fromFile` for a Java API + * @see `Readers.fromStream` for a Java API */ - @throws(classOf[FileNotFoundException]) - @throws(classOf[SecurityException]) - def fromFile(f: File): Reader[Buf] = - fromStream(new FileInputStream(f)) + def fromStream(s: InputStream, chunkSize: Int): Reader[Buf] = InputStreamReader(s, chunkSize) /** - * Wrap `InputStream` with a [[Reader]]. + * Create a new [[Reader]] from a given `Seq`. * - * Note that the given `InputStream` will be closed - * on reading of EOF and [[Reader.discard()]]. + * The resources held by the returned [[Reader]] are released on reading of EOF and + * [[Reader.discard]]. * - * @see `Readers.fromStream` for a Java API + * @note Multiple outstanding reads are not allowed on this reader. */ - def fromStream(s: InputStream): Reader[Buf] = - InputStreamReader(s) + def fromSeq[A](seq: Seq[A]): Reader[A] = fromIterator(seq.iterator) + + /** + * Java-Friendly version of fromSeq + * Create a new [[Reader]] from a given Java `List` + */ + def fromCollection[A](list: JCollection[A]): Reader[A] = fromSeq( + list.iterator.asScala.toSeq + ) /** - * Allow [[AsyncStream]] to be consumed as a [[Reader]] + * Allow [[com.twitter.concurrent.AsyncStream]] to be consumed as a [[Reader]] */ - def fromAsyncStream[A <: Buf](as: AsyncStream[A]): Reader[A] = { - val pipe = new Pipe[A]() + def fromAsyncStream[A](as: AsyncStream[A]): Reader[A] = { + val pipe = new Pipe[A] // orphan the Future but allow it to clean up // the Pipe IF the stream ever finishes or fails as.foreachF(pipe.write).respond { @@ -209,96 +335,117 @@ object Reader { } /** - * Convenient abstraction to read from a stream of Readers as if it were a - * single Reader. + * Transformation (or lift) from [[Reader]] into `AsyncStream`. */ - def concat(readers: AsyncStream[Reader[Buf]]): Reader[Buf] = { - val target = new Pipe[Buf]() - val f = copyMany(readers, target).respond { + def toAsyncStream[A](r: Reader[A]): AsyncStream[A] = + AsyncStream.fromFuture(r.read()).flatMap { + case Some(buf) => buf +:: Reader.toAsyncStream(r) + case None => AsyncStream.empty[A] + } + + /** + * Convenient abstraction to read from a stream (AsyncStream) of Readers as if + * it were a single Reader. + * @param readers An AsyncStream holds a stream of Reader[A] + */ + def concat[A](readers: AsyncStream[Reader[A]]): Reader[A] = { + val target = new Pipe[A] + val copied = readers.foreachF(Pipe.copy(_, target)).respond { case Throw(exc) => target.fail(exc) case _ => target.close() } - new Reader[Buf] { - def read(n: Int): Future[Option[Buf]] = target.read(n) - def discard(): Unit = { + + target.onClose.respond { + case Return(StreamTermination.Discarded) => // We have to do this so that when the the target is discarded we can // interrupt the read operation. Consider the following: // // r.read(..) { case Some(b) => target.write(b) } // - // The computation r.read(..) will be interupted because we set an + // The computation r.read(..) will be interrupted because we set an // interrupt handler in Reader.copy to discard `r`. - f.raise(new Reader.ReaderDiscarded()) - target.discard() - } + copied.raise(new ReaderDiscardedException()) + case _ => () } - } - /** - * Copy bytes from many Readers to a Writer. The Writer is unmanaged, the - * caller is responsible for finalization and error handling, e.g.: - * - * {{{ - * Reader.copyMany(readers, writer) ensure writer.close() - * }}} - * - * @param bufsize The number of bytes to read each time. - */ - def copyMany(readers: AsyncStream[Reader[Buf]], target: Writer[Buf], bufsize: Int): Future[Unit] = - readers.foreachF(Reader.copy(_, target, bufsize)) + target + } /** - * Copy bytes from many Readers to a Writer. The Writer is unmanaged, the - * caller is responsible for finalization and error handling, e.g.: - * - * {{{ - * Reader.copyMany(readers, writer) ensure writer.close() - * }}} + * Convenient abstraction to read from a collection of Readers as if + * it were a single Reader. + * @param readers A collection of Reader[A] */ - def copyMany(readers: AsyncStream[Reader[Buf]], target: Writer[Buf]): Future[Unit] = - copyMany(readers, target, Writer.BufferSize) + def concat[A](readers: Seq[Reader[A]]): Reader[A] = { + Reader.fromSeq(readers).flatten + } /** - * Copy the bytes from a Reader to a Writer in chunks of size `n`. The Writer - * is unmanaged, the caller is responsible for finalization and error - * handling, e.g.: + * Convenient abstraction to read from a stream (Reader) of Readers as if it were a single Reader. + * @param readers A Reader holds a stream of Reader[A] * - * {{{ - * Reader.copy(r, w, n) ensure w.close() - * }}} - * - * @param n The number of bytes to read on each refill of the Writer. + * @note All operations of the new Reader will be in sync with the outermost Reader. Discarding + * one Reader will discard the other Reader. When one Reader's onClose resolves, the other + * Reader's onClose will be resolved immediately with the same value. + * The subsequent readers are unmanaged, the caller is responsible for discarding those + * when abandoned. */ - def copy(r: Reader[Buf], w: Writer[Buf], n: Int): Future[Unit] = { - def loop(): Future[Unit] = - r.read(n).flatMap { - case None => Future.Done - case Some(buf) => w.write(buf) before loop() + def flatten[A](readers: Reader[Reader[A]]): Reader[A] = new Reader[A] { self => + // access currentReader and curReaderClosep are synchronized on `self` + private[this] var currentReader: Reader[A] = Reader.empty + private[this] var curReaderClosep: Promise[StreamTermination] with Detachable = + Promise.attached(currentReader.onClose) + + private[this] val closep = Promise[StreamTermination]() + + readers.onClose.respond(closep.updateIfEmpty) + + def read(): Future[Option[A]] = { + self.synchronized(currentReader).read().transform { + case Return(None) => + updateCurrentAndRead() + case Return(sa) => + Future.value(sa) + case t @ Throw(_) => + // update `closep` with the exception thrown before discarding readers, + // because discard will update `closep` to a `Discarded` StreamTermination + // and return a `ReaderDiscardedException` during reading. + closep.updateIfEmpty(t.cast[StreamTermination]) + readers.discard() + Future.const(t) } - val p = new Promise[Unit] - // We have to do this because discarding the writer doesn't interrupt read - // operations, it only fails the next write operation. - loop().proxyTo(p) - p.setInterruptHandler { case exc => r.discard() } - p - } + } - /** - * Copy the bytes from a Reader to a Writer in chunks of size - * `Writer.BufferSize`. The Writer is unmanaged, the caller is responsible - * for finalization and error handling, e.g.: - * - * {{{ - * Reader.copy(r, w) ensure w.close() - * }}} - */ - def copy(r: Reader[Buf], w: Writer[Buf]): Future[Unit] = copy(r, w, Writer.BufferSize) + def discard(): Unit = { + self + .synchronized { + curReaderClosep.detach() + currentReader + }.discard() + readers.discard() + } - /** - * Wraps a [[ Reader[Buf] ]] and emits frames as decided by `framer`. - * - * @note The returned `Reader` may not be thread safe depending on the behavior - * of the framer. - */ - def framed(r: Reader[Buf], framer: Buf => Seq[Buf]): Reader[Buf] = new Framed(r, framer) + def onClose: Future[StreamTermination] = closep + + /** + * Update currentReader and the callback promise to be the `closep` of currentReader + */ + private def updateCurrentAndRead(): Future[Option[A]] = { + readers.read().flatMap { + case Some(reader) => + val curClosep = self.synchronized { + currentReader = reader + currentReader.onClose + } + curReaderClosep = Promise.attached(curClosep) + curReaderClosep.respond { + case Return(StreamTermination.FullyRead) => + case discardedOrFailed => closep.updateIfEmpty(discardedOrFailed) + } + read() + case _ => + Future.None + } + } + } } diff --git a/util-core/src/main/scala/com/twitter/io/StreamTermination.scala b/util-core/src/main/scala/com/twitter/io/StreamTermination.scala new file mode 100644 index 0000000000..fddc326245 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/StreamTermination.scala @@ -0,0 +1,30 @@ +package com.twitter.io + +/** + * Trait that indicates what kind of stream termination a stream saw. + */ +sealed abstract class StreamTermination { + def isFullyRead: Boolean +} + +object StreamTermination { + + /** + * Indicates a stream terminated when it was read all the way until the end. + */ + case object FullyRead extends StreamTermination { + val Return: com.twitter.util.Return[FullyRead.type] = com.twitter.util.Return(this) + + def isFullyRead: Boolean = true + } + + /** + * Indicates a stream terminated when it was discarded by the reader before it + * could be read all the way until the end. + */ + case object Discarded extends StreamTermination { + val Return: com.twitter.util.Return[Discarded.type] = com.twitter.util.Return(this) + + def isFullyRead: Boolean = false + } +} diff --git a/util-core/src/main/scala/com/twitter/io/TempFile.scala b/util-core/src/main/scala/com/twitter/io/TempFile.scala index 0bc6908f76..476d995e83 100644 --- a/util-core/src/main/scala/com/twitter/io/TempFile.scala +++ b/util-core/src/main/scala/com/twitter/io/TempFile.scala @@ -10,7 +10,7 @@ object TempFile { * * Note, due to the usage of `File.deleteOnExit()` callers should * be careful using this as it can leak memory. - * See http://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for + * See https://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for * example. * * @param path the resource-relative path to make a temp file from @@ -24,7 +24,7 @@ object TempFile { * * Note, due to the usage of `File.deleteOnExit()` callers should * be careful using this as it can leak memory. - * See http://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for + * See https://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for * example. * * @param klass the `Class` to use for getting the resource. @@ -41,7 +41,7 @@ object TempFile { * * Note, due to the usage of `File.deleteOnExit()` callers should * be careful using this as it can leak memory. - * See http://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for + * See https://bugs.sun.com/bugdatabase/view_bug.do?bug_id=4513817 for * example. * * @param path the path to make a temp file from @@ -55,7 +55,12 @@ object TempFile { case null => throw new FileNotFoundException(path) case stream => - val (basename, ext) = parsePath(path) + val (basename, ext) = parsePath(path) match { + // The prefix parameter of File.createTempFile(prefix,suffix) is + // expected to be 3 characters long. + case (basename, ext) if basename.length < 3 => (basename + "tmp", ext) + case basenameExt => basenameExt + } val file = File.createTempFile(basename, "." + ext) file.deleteOnExit() val fos = new BufferedOutputStream(new FileOutputStream(file), 1 << 20) diff --git a/util-core/src/main/scala/com/twitter/util/TempFolder.scala b/util-core/src/main/scala/com/twitter/io/TempFolder.scala similarity index 74% rename from util-core/src/main/scala/com/twitter/util/TempFolder.scala rename to util-core/src/main/scala/com/twitter/io/TempFolder.scala index 28174fdb5b..36eeed4fb9 100644 --- a/util-core/src/main/scala/com/twitter/util/TempFolder.scala +++ b/util-core/src/main/scala/com/twitter/io/TempFolder.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -14,23 +14,19 @@ * limitations under the License. */ -package com.twitter.util +package com.twitter.io import java.io.File -import com.twitter.io.Files - /** * Test mixin that creates a temporary thread-local folder for a block of code to execute in. * The folder is recursively deleted after the test. * - * Note, the [[com.twitter.util.io]] package would be a better home for this trait. - * * Note that multiple uses of TempFolder cannot be nested, because the temporary directory * is effectively a thread-local global. */ trait TempFolder { - private val _folderName = new ThreadLocal[File] + private val _folderName: ThreadLocal[File] = new ThreadLocal[File] /** * Runs the given block of code with the presence of a temporary folder whose name can be @@ -42,10 +38,10 @@ trait TempFolder { val tempFolder = System.getProperty("java.io.tmpdir") // Note: If we were willing to have a dependency on Guava in util-core // we could just use `com.google.common.io.Files.createTempDir()` - var folder: File = null - do { + var folder: File = new File(tempFolder, "scala-test-" + System.currentTimeMillis) + while (!folder.mkdir()) { folder = new File(tempFolder, "scala-test-" + System.currentTimeMillis) - } while (!folder.mkdir()) + } _folderName.set(folder) try { @@ -57,13 +53,13 @@ trait TempFolder { /** * @return The current thread's active temporary folder. - * @throws RuntimeException if not running within a withTempFolder block + * @note Throws `RuntimeException` if not running within a withTempFolder block */ - def folderName = { _folderName.get.getPath } + def folderName: String = { _folderName.get.getPath } /** * @return The canonical path of the current thread's active temporary folder. - * @throws RuntimeException if not running within a withTempFolder block + * @note Throws `RuntimeException` if not running within a withTempFolder block */ - def canonicalFolderName = { _folderName.get.getCanonicalPath } + def canonicalFolderName: String = { _folderName.get.getCanonicalPath } } diff --git a/util-core/src/main/scala/com/twitter/io/Utf8Bytes.scala b/util-core/src/main/scala/com/twitter/io/Utf8Bytes.scala new file mode 100644 index 0000000000..060f7afb67 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/io/Utf8Bytes.scala @@ -0,0 +1,77 @@ +package com.twitter.io + +import java.nio.charset.StandardCharsets + +object Utf8Bytes { + + /** + * Create a [[Utf8Bytes]] instance from a byte array. The array passed in is copied to ensure it can not be later + * mutated. + * + * @note No validation is done on the passed bytes to ensure that they are valid UTF-8. If invalid bytes are present + * they will be converted to the default UTF-8 replacement string. + * @note As an optimization, if it's known that the byte[] will never be mutated once passed to the [[Utf8Bytes]] + * instance, use `Utf8Bytes(Buf.ByteArray.Owned(bytes))` to create an instance without copying the byte[]. + */ + def apply(bytes: Array[Byte]): Utf8Bytes = apply(Buf.ByteArray.Shared(bytes)) + + /** + * Create a [[Utf8Bytes]] instance from a [[String]]. The string is first encoded into UTF-8 bytes. + */ + def apply(str: String): Utf8Bytes = apply(Buf.Utf8(str)) + + /** + * Create a [[Utf8Bytes]] instance from a [[Buf]]. + * + * @note No validation is done on the passed bytes to ensure that they are valid UTF-8. If invalid bytes are present + * they will be converted to the default UTF-8 replacement string. + * @note This is the most efficient way to create an instance. It neither has to encode a string to bytes nor do a + * defensive copy, since [[Buf]] instances are already immutable. + */ + def apply(buf: Buf): Utf8Bytes = new Utf8Bytes(buf) + + val empty: Utf8Bytes = apply(Buf.Empty) +} + +/** + * An immutable representation of a UTF-8 encoded string. The representation is lazily converted to a [[String]] when + * operations require it. + * + * @note This class is thread-safe. + */ +final class Utf8Bytes private ( + private[this] val buf: Buf) + extends CharSequence { + + private[this] var str: String = _ + + /** + * The underlying byte content of this instance. + */ + def asBuf: Buf = buf + + override def toString: String = { + // multiple threads may race to set `str`, however doing so is safe + if (str == null) { + str = Buf.decodeString(buf, StandardCharsets.UTF_8) + } + str + } + + override def equals(obj: Any): Boolean = + obj match { + case other: AnyRef if this.eq(other) => true + case other: Utf8Bytes => buf == other.asBuf + case _ => false + } + + override def hashCode(): Int = buf.hashCode + + // note, all methods below must operate on chars, not bytes, so they must operate on the string representation + // rather than the byte representation + override def length(): Int = toString.length + + override def charAt(index: Int): Char = toString.charAt(index) + + override def subSequence(start: Int, end: Int): CharSequence = toString.subSequence(start, end) +} diff --git a/util-core/src/main/scala/com/twitter/io/Writer.scala b/util-core/src/main/scala/com/twitter/io/Writer.scala index 9c37d688ba..ae4d465dcc 100644 --- a/util-core/src/main/scala/com/twitter/io/Writer.scala +++ b/util-core/src/main/scala/com/twitter/io/Writer.scala @@ -4,54 +4,113 @@ import com.twitter.util.{Closable, Future, Time} import java.io.OutputStream /** - * A Writer represents a sink for a stream of `A`s, providing a convenient interface - * for the producer of such streams. + * A writer represents a sink for a stream of arbitrary elements. While [[Reader]] provides an API + * to consume streams, [[Writer]] provides an API to produce streams. + * + * Similar to [[Reader]]s, a very typical way to work with [[Writer]]s is to model a stream producer + * as a write-loop: + * + * {{{ + * def produce[A](w: Writer[A])(generate: () => Option[A]): Future[Unit] = + * generate match { + * case Some(a) => w.write(a).before(produce(w)(generate)) + * case None => w.close() // signal EOF and quit producing + * } + * }}} + * + * == Closures and Failures == + * + * Writers are [[com.twitter.util.Closable]] and the producer MUST `close` the stream when it + * finishes or `fail` the stream when it encounters a failure and can no longer continue. Streams + * backed by network connections are particularly prone to resource leaks when they aren't + * cleaned up properly. + * + * Observing stream failures in the producer could be done via a [[com.twitter.util.Future]] + * returned form a write-loop: + * + * {{{ + * produce(writer)(generator).respond { + * case Return(()) => println("Produced a stream successfully.") + * case Throw(e) => println(s"Could not produce a stream because of a failure: \$e") + * } + * }}} + * + * @note Encountering a stream failure would terminate the write-loop given the + * [[com.twitter.util.Future]] recursion semantics. + * + * @note Once failed or closed, a stream can not be restarted such that all future writes will + * resolve into a failure. + * + * @note Closing an already failed stream does not have an effect. + * + * == Back Pressure == + * + * By analogy with read-loops (see [[Reader]] API), write-loops leverage[[com.twitter.util.Future]] + * recursion to exert back-pressure: the next write isn't issued until the previous write finishes. + * This will always ensure a finer grained back-pressure in network systems allowing the producers + * to adjust the flow rate based on consumer's speed and not on IO buffering. + * + * @note Whether or not multiple pending writes are allowed on a `Writer` type is an undefined + * behaviour but could be changed in a refinement. */ -trait Writer[-A] { +trait Writer[-A] extends Closable { self => /** - * Write a chunk. The returned future is completed when the chunk has been - * fully read by the sink. + * Write an `element` into this stream. Although undefined by this contract, a trustworthy + * implementation (such as [[Pipe]]) would do its best to resolve the returned + * [[com.twitter.util.Future]] only when a consumer observes a written `element`. * - * Only one outstanding `write` is permitted, thus backpressure is asserted. + * The write can also resolve into a failure (failed [[com.twitter.util.Future]]). */ - def write(buf: A): Future[Unit] + def write(element: A): Future[Unit] /** - * Indicate that the producer of the bytestream has failed. No further writes are allowed. + * Fail this stream with a given `cause`. No further writes are allowed, but if happen, will + * resolve into a [[com.twitter.util.Future]] failed with `cause`. + * + * @note Failing an already closed stream does not have an effect. */ def fail(cause: Throwable): Unit -} - -/** - * @see Writers for Java friendly APIs. - */ -object Writer { /** - * A [[ClosableWriter]] instance that will always fail. Useful for situations - * where writing to the [[Writer]] is nonsensical such as the [[Writer]] on the - * `Response` returned by the Finagle HTTP client. + * A [[com.twitter.util.Future]] that resolves once this writer is closed and flushed. + * + * It may contain an error if the writer is failed. + * + * This is useful for any extra resource cleanup that you must do once the stream is no longer + * being used. */ - val FailingWriter: Writer[Buf] with Closable = fail[Buf] - - def fail[A]: Writer[A] with Closable = new Writer[A] with Closable { - def write(buf: A): Future[Unit] = Future.exception(new IllegalStateException("NullWriter")) - def fail(cause: Throwable): Unit = () - def close(deadline: Time): Future[Unit] = Future.Done - } + def onClose: Future[StreamTermination] /** - * Represents a [[Writer]] which is [[Closable]]. - * - * Exists primarily for Java compatibility. + * Given f, a function from B into A, creates an Writer[B] whose `fail` and `close` functions + * are equivalent to Writer[A]'s. Writer[B]'s `write` function is equivalent to: + * {{{ + * def write(element: B) = Writer[A].write(f(element)) + * }}} */ - trait ClosableWriter[A] extends Writer[A] with Closable + final def contramap[B](f: B => A): Writer[B] = new Writer[B] { + def write(element: B): Future[Unit] = self.write(f(element)) + def fail(cause: Throwable): Unit = self.fail(cause) + def close(deadline: Time): Future[Unit] = self.close(deadline) + def onClose: Future[StreamTermination] = self.onClose + } +} + +/** + * Abstract `Writer` class for Java compatibility. + */ +abstract class AbstractWriter[-A] extends Writer[A] + +/** + * @see Writers for Java friendly APIs. + */ +object Writer { val BufferSize: Int = 4096 /** - * Construct a [[ClosableWriter]] from a given OutputStream. + * Construct a [[Writer]] from a given OutputStream. * * This [[Writer]] is not thread safe. If multiple threads attempt to `write`, the * behavior is identical to multiple threads calling `write` on the underlying @@ -59,16 +118,16 @@ object Writer { * * @param bufsize Size of the copy buffer between Writer and OutputStream. */ - def fromOutputStream(out: OutputStream, bufsize: Int): ClosableWriter[Buf] = + def fromOutputStream(out: OutputStream, bufsize: Int): Writer[Buf] = new OutputStreamWriter(out, bufsize) /** - * Construct a [[ClosableWriter]] from a given OutputStream. + * Construct a [[Writer]] from a given OutputStream. * * This [[Writer]] is not thread safe. If multiple threads attempt to `write`, the * behavior is identical to multiple threads calling `write` on the underlying * OutputStream. */ - def fromOutputStream(out: OutputStream): ClosableWriter[Buf] = + def fromOutputStream(out: OutputStream): Writer[Buf] = fromOutputStream(out, BufferSize) } diff --git a/util-core/src/main/scala/com/twitter/io/exp/ActivitySource.scala b/util-core/src/main/scala/com/twitter/io/exp/ActivitySource.scala deleted file mode 100644 index db9e4a6ffa..0000000000 --- a/util-core/src/main/scala/com/twitter/io/exp/ActivitySource.scala +++ /dev/null @@ -1,204 +0,0 @@ -package com.twitter.io.exp - -import com.twitter.conversions.time._ -import com.twitter.io.{InputStreamReader, Buf, Reader} -import com.twitter.util._ -import java.io.{FileInputStream, File} -import java.lang.ref.{ReferenceQueue, WeakReference} -import java.util.concurrent.atomic.AtomicBoolean -import java.util.HashMap - -/** - * An ActivitySource provides access to observerable named variables. - */ -trait ActivitySource[+T] { - - /** - * Returns an [[com.twitter.util.Activity]] for a named T-typed variable. - */ - def get(name: String): Activity[T] - - /** - * Produces an ActivitySource which queries this ActivitySource, falling back to - * a secondary ActivitySource only when the primary result is ActivitySource.Failed. - * - * @param that the secondary ActivitySource - */ - def orElse[U >: T](that: ActivitySource[U]): ActivitySource[U] = - new ActivitySource.OrElse(this, that) -} - -object ActivitySource { - - /** - * A Singleton exception to indicate that an ActivitySource failed to find - * a named variable. - */ - object NotFound extends IllegalStateException - - /** - * An ActivitySource for observing file contents. Once observed, - * each file will be polled once per period. - */ - def forFiles(period: Duration = 1.minute)(implicit timer: Timer): ActivitySource[Buf] = - new CachingActivitySource(new FilePollingActivitySource(period)(timer)) - - /** - * Create an ActivitySource for ClassLoader resources. - */ - def forClassLoaderResources( - cl: ClassLoader = ClassLoader.getSystemClassLoader - ): ActivitySource[Buf] = - new CachingActivitySource(new ClassLoaderActivitySource(cl)) - - private[ActivitySource] class OrElse[T, U >: T]( - primary: ActivitySource[T], - failover: ActivitySource[U] - ) extends ActivitySource[U] { - def get(name: String): Activity[U] = { - primary.get(name) transform { - case Activity.Failed(_) => failover.get(name) - case state => Activity(Var.value(state)) - } - } - } -} - -/** - * A convenient wrapper for caching the results returned by the - * underlying ActivitySource. - */ -class CachingActivitySource[T](underlying: ActivitySource[T]) extends ActivitySource[T] { - - private[this] val refq = new ReferenceQueue[Activity[T]] - private[this] val forward = new HashMap[String, WeakReference[Activity[T]]] - private[this] val reverse = new HashMap[WeakReference[Activity[T]], String] - - /** - * A caching proxy to the underlying ActivitySource. Vars are cached by - * name, and are tracked with WeakReferences. - */ - def get(name: String): Activity[T] = synchronized { - gc() - Option(forward.get(name)) flatMap { wr => - Option(wr.get()) - } match { - case Some(v) => v - case None => - val v = underlying.get(name) - val ref = new WeakReference(v, refq) - forward.put(name, ref) - reverse.put(ref, name) - v - } - } - - /** - * Remove garbage collected cache entries. - */ - def gc(): Unit = synchronized { - var ref = refq.poll() - while (ref != null) { - val key = reverse.remove(ref) - if (key != null) - forward.remove(key) - - ref = refq.poll() - } - } -} - -/** - * An ActivitySource for observing the contents of a file with periodic polling. - */ -class FilePollingActivitySource private[exp] ( - period: Duration, - pool: FuturePool -)(implicit timer: Timer) - extends ActivitySource[Buf] { - - private[exp] def this(period: Duration)(implicit timer: Timer) = - this(period, FuturePool.unboundedPool) - - import com.twitter.io.exp.ActivitySource._ - - def get(name: String): Activity[Buf] = { - val v = Var.async[Activity.State[Buf]](Activity.Pending) { value => - val timerTask = timer.schedule(Time.now, period) { - val file = new File(name) - - if (file.exists()) { - pool { - val reader = new InputStreamReader( - new FileInputStream(file), - InputStreamReader.DefaultMaxBufferSize, - pool - ) - Reader.readAll(reader) respond { - case Return(buf) => - value() = Activity.Ok(buf) - case Throw(cause) => - value() = Activity.Failed(cause) - } ensure { - // InputStreamReader ignores the deadline in close - reader.close(Time.Undefined) - } - } - } else { - value() = Activity.Failed(NotFound) - } - } - - Closable.make { _ => - Future { timerTask.cancel() } - } - } - - Activity(v) - } -} - -/** - * An ActivitySource for ClassLoader resources. - */ -class ClassLoaderActivitySource private[exp] (classLoader: ClassLoader, pool: FuturePool) - extends ActivitySource[Buf] { - - import com.twitter.io.exp.ActivitySource._ - - private[exp] def this(classLoader: ClassLoader) = this(classLoader, FuturePool.unboundedPool) - - def get(name: String): Activity[Buf] = { - // This Var is updated at most once since ClassLoader - // resources don't change (do they?). - val runOnce = new AtomicBoolean(false) - val p = new Promise[Activity.State[Buf]] - - // Defer loading until the first observation - val v = Var.async[Activity.State[Buf]](Activity.Pending) { value => - if (runOnce.compareAndSet(false, true)) { - pool { - classLoader.getResourceAsStream(name) match { - case null => p.setValue(Activity.Failed(NotFound)) - case stream => - val reader = - new InputStreamReader(stream, InputStreamReader.DefaultMaxBufferSize, pool) - Reader.readAll(reader) respond { - case Return(buf) => - p.setValue(Activity.Ok(buf)) - case Throw(cause) => - p.setValue(Activity.Failed(cause)) - } ensure { - // InputStreamReader ignores the deadline in close - reader.close(Time.Undefined) - } - } - } - } - p.onSuccess(value() = _) - Closable.nop - } - - Activity(v) - } -} diff --git a/util-core/src/main/scala/com/twitter/io/exp/BelowThroughputException.scala b/util-core/src/main/scala/com/twitter/io/exp/BelowThroughputException.scala deleted file mode 100644 index 3667a54489..0000000000 --- a/util-core/src/main/scala/com/twitter/io/exp/BelowThroughputException.scala +++ /dev/null @@ -1,15 +0,0 @@ -package com.twitter.io.exp - -import com.twitter.util.{Duration, TimeoutException} - -/** - * Indicates a [[com.twitter.io.Reader reader]] or - * [[com.twitter.io.Writer writer]] dropped below - * the minimum required bytes per second threshold. - * - * @param elapsed total time spent reading or writing. - * - * @see [[MinimumThroughput]] - */ -case class BelowThroughputException(elapsed: Duration, currentBps: Double, expectedBps: Double) - extends TimeoutException(s"current bps $currentBps below minimum $expectedBps") diff --git a/util-core/src/main/scala/com/twitter/io/exp/MinimumThroughput.scala b/util-core/src/main/scala/com/twitter/io/exp/MinimumThroughput.scala deleted file mode 100644 index a756346931..0000000000 --- a/util-core/src/main/scala/com/twitter/io/exp/MinimumThroughput.scala +++ /dev/null @@ -1,158 +0,0 @@ -package com.twitter.io.exp - -import com.twitter.conversions.time._ -import com.twitter.io.{Writer, Buf, Reader} -import com.twitter.util._ -import scala.util.control.NoStackTrace - -/** - * [[Reader Readers]] and [[Writer Writers]] that have minimum bytes - * per second throughput limits. - */ -object MinimumThroughput { - - private val BpsNoElapsed = -1d - - private val MinDeadline = 1.second - - /** internal marker used to raise on read/write timeouts. */ - private case object MinThroughputTimeoutException extends Exception with NoStackTrace - - private abstract class Min(minBps: Double) { - assert(minBps >= 0, s"minBps must be non-negative: $minBps") - - // we rely on the `reader.read` and `writer.write` contract that only - // one read or write can ever be outstanding, so volatile provides - // us thread-safety. - @volatile protected[this] var bytes = 0L - @volatile protected[this] var elapsed = Duration.Zero - - /** Calculate and return the current bps */ - def bps: Double = { - val elapsedSecs = elapsed.inSeconds - if (elapsedSecs == 0) { - BpsNoElapsed - } else { - bytes.toDouble / elapsedSecs.toDouble - } - } - - /** - * Return a cap on how long we're willing to wait for the operation to - * complete processing `numBytes` while staying above the min bps requirements. - */ - def nextDeadline(numBytes: Int): Duration = - if (minBps == 0) - Duration.Top - else - Duration.fromSeconds(((bytes + numBytes) / minBps).toInt).max(MinDeadline) - - /** Verify if we are over the min bps requirements. */ - def throughputOk: Boolean = { - val newBps = bps - newBps >= minBps || newBps == BpsNoElapsed - } - - def newBelowThroughputException: BelowThroughputException = - new BelowThroughputException(elapsed, bps, minBps) - - } - - private class MinReader(reader: Reader[Buf], minBps: Double, timer: Timer) - extends Min(minBps) - with Reader[Buf] { - def discard(): Unit = reader.discard() - - def read(n: Int): Future[Option[Buf]] = { - val deadline = nextDeadline(n) - - val start = Time.now - val read: Future[Option[Buf]] = - reader.read(n).raiseWithin(timer, deadline, MinThroughputTimeoutException) - - read.transform { res => - elapsed += Time.now - start - res match { - case Throw(MinThroughputTimeoutException) => - discard() - Future.exception(newBelowThroughputException) - case Throw(_) => - read // pass through other failures untouched - case Return(None) => - read // hit EOF, so we're fine - case Return(Some(buf)) => - // read some bytes, see if we're still good in terms of bps - bytes += buf.length - if (throughputOk) { - read - } else { - discard() - Future.exception(newBelowThroughputException) - } - } - } - } - } - - private class MinWriter(writer: Writer[Buf], minBps: Double, timer: Timer) - extends Min(minBps) - with Writer[Buf] { - def fail(cause: Throwable): Unit = writer.fail(cause) - - override def write(buf: Buf): Future[Unit] = { - val numBytes = buf.length - val deadline = nextDeadline(numBytes) - - val start = Time.now - val write: Future[Unit] = - writer.write(buf).raiseWithin(timer, deadline, MinThroughputTimeoutException) - - write.transform { res => - elapsed += Time.now - start - res match { - case Throw(MinThroughputTimeoutException) => - val ex = newBelowThroughputException - fail(ex) - Future.exception(ex) - case Throw(_) => - write // pass through other failures untouched - case Return(_) => - // wrote some bytes, see if we're still good in terms of bps - bytes += numBytes - if (throughputOk) { - write - } else { - val ex = newBelowThroughputException - fail(ex) - Future.exception(ex) - } - } - } - } - } - - /** - * Ensures that the throughput of [[Reader.read reads]] happen - * above a minimum specified bytes per second rate. - * - * Reads will return a Future of a [[BelowThroughputException]] when the - * throughput is not met. - * - * @note the minimum time allotted to any read is 1 second. - */ - def reader(reader: Reader[Buf], minBps: Double, timer: Timer): Reader[Buf] = - new MinReader(reader, minBps, timer) - - /** - * Ensures that the throughput of [[Writer.write writes]] happen - * above a minimum specified bytes per second rate. - * - * Writes will return a Future of a [[BelowThroughputException]] when the - * throughput is not met. - * - * @note the minimum time allotted to any write is 1 second. - */ - def writer(writer: Writer[Buf], minBps: Double, timer: Timer): Writer[Buf] = - new MinWriter(writer, minBps, timer) - -} diff --git a/util-core/src/main/scala/com/twitter/util/Activity.scala b/util-core/src/main/scala/com/twitter/util/Activity.scala index 04eb57d687..5fbadd9a3a 100644 --- a/util-core/src/main/scala/com/twitter/util/Activity.scala +++ b/util-core/src/main/scala/com/twitter/util/Activity.scala @@ -1,25 +1,46 @@ package com.twitter.util import java.util.{List => JList} - -import scala.collection.generic.CanBuild -import scala.collection.JavaConverters._ +import scala.collection.mutable.ArrayBuffer import scala.collection.mutable.Buffer -import scala.language.higherKinds -import scala.reflect.ClassTag +import scala.jdk.CollectionConverters._ import scala.util.control.NonFatal +import scala.collection.{Seq => AnySeq} /** - * An Activity is a handle to a concurrently running process, producing - * T-typed values. An activity is in one of three states: + * Activity, like [[com.twitter.util.Var Var]], is a value that varies over + * time; but, richer than Var, these values have error and pending states. + * They are useful for modeling a process that produces values, but also + * intermittent failures (e.g. address resolution, reloadable configuration). + * + * As such, an Activity can be in one of three states: * * - [[com.twitter.util.Activity.Pending Pending]]: output is pending; * - [[com.twitter.util.Activity.Ok Ok]]: an output is available; and * - [[com.twitter.util.Activity.Failed Failed]]: the process failed with an exception. * - * An activity may transition between any state at any time. + * An Activity may transition between any state at any time. They are typically + * constructed from [[com.twitter.util.Event Events]], but they can also be + * constructed without arguments. + * + * {{{ + * val (activity, witness) = Activity[Int]() + * println(activity.sample()) // throws java.lang.IllegalStateException: Still pending + * }}} + * + * As you can see, Activities start in a Pending state. Attempts to sample an + * Activity that is pending will fail. Curiously, the constructor for an + * Activity also returns a [[com.twitter.util.Witness Witness]]. + * + * Recall that a Witness is the interface to something that you can send values + * to; and in this case, that something is the Activity. * - * (The observant reader will notice that this is really a monad + * {{{ + * witness.notify(Return(1)) + * println(activity.sample()) // prints 1 + * }}} + * + * (For implementers, it may be useful to think of Activity as a monad * transformer for an ''Op'' monad over [[com.twitter.util.Var Var]] * where Op is like a [[com.twitter.util.Try Try]] with an additional * pending state.) @@ -63,10 +84,11 @@ case class Activity[+T](run: Var[Activity.State[T]]) { def flatMap[U](f: T => Activity[U]): Activity[U] = Activity(run flatMap { case Ok(v) => - val a = try f(v) - catch { - case NonFatal(exc) => Activity.exception(exc) - } + val a = + try f(v) + catch { + case NonFatal(exc) => Activity.exception(exc) + } a.run case Pending => Var.value(Activity.Pending) @@ -79,10 +101,11 @@ case class Activity[+T](run: Var[Activity.State[T]]) { */ def transform[U](f: Activity.State[T] => Activity[U]): Activity[U] = Activity(run flatMap { act => - val a = try f(act) - catch { - case NonFatal(exc) => Activity.exception(exc) - } + val a = + try f(act) + catch { + case NonFatal(exc) => Activity.exception(exc) + } a.run }) @@ -121,12 +144,17 @@ case class Activity[+T](run: Var[Activity.State[T]]) { * once an [[com.twitter.util.Activity.Ok Ok]] value has been returned * it stops error propagation and instead returns the last Ok value. */ - def stabilize: Activity[T] = - Activity(states.foldLeft[Activity.State[T]](Activity.Pending) { - case (_, next @ Activity.Ok(_)) => next - case (prev @ Activity.Ok(_), _) => prev - case (_, next) => next + def stabilize: Activity[T] = { + @volatile var stable: Activity.State[T] = Activity.Pending + Activity(run.map { + case next @ Activity.Ok(_) => { + stable = next + next + } + case other if stable == Activity.Pending => other + case other => stable }) + } } /** @@ -152,19 +180,19 @@ object Activity { * Constructs an Activity from a state Event. */ def apply[T](states: Event[State[T]]): Activity[T] = - Activity(Var(Pending, states)) + Activity(Var.async[State[T]](Pending) { updatable => + states.respond { state => + updatable() = state + } + }) /** * Collect a collection of activities into an activity of a collection * of values. * - * @usecase def collect[T](activities: Coll[Activity[T]]): Activity[Coll[T]] + * @example def collect[T](activities: Coll[Activity[T]]): Activity[Coll[T]] */ - def collect[T: ClassTag, CC[X] <: Traversable[X]]( - acts: CC[Activity[T]] - )( - implicit newBuilder: CanBuild[T, CC[T]] - ): Activity[CC[T]] = { + def collect[T](acts: AnySeq[Activity[T]]): Activity[Seq[T]] = { collect(acts, false) } @@ -176,34 +204,28 @@ object Activity { * Like Var.collectIndependent this is a workaround and should be deprecated when a version * of Var.collect without a stack overflow issue is implemented. * - * @usecase def collectIndependent[T](activities: Coll[Activity[T]]): Activity[Coll[T]] + * @example def collectIndependent[T](activities: Coll[Activity[T]]): Activity[Coll[T]] */ - def collectIndependent[T: ClassTag, CC[X] <: Traversable[X]]( - acts: CC[Activity[T]] - )( - implicit newBuilder: CanBuild[T, CC[T]] - ): Activity[CC[T]] = { + def collectIndependent[T](acts: AnySeq[Activity[T]]): Activity[Seq[T]] = { collect(acts, true) } - private[this] def collect[T: ClassTag, CC[X] <: Traversable[X]]( - acts: CC[Activity[T]], + private[this] def collect[T]( + acts: AnySeq[Activity[T]], collectIndependent: Boolean - )( - implicit newBuilder: CanBuild[T, CC[T]] - ): Activity[CC[T]] = { + ): Activity[Seq[T]] = { if (acts.isEmpty) - return Activity.value(newBuilder().result) + return Activity.value(Seq.empty[T]) - val states: Traversable[Var[State[T]]] = acts.map(_.run) - val stateVar: Var[Traversable[State[T]]] = if (collectIndependent) { + val states: AnySeq[Var[State[T]]] = acts.map(_.run) + val stateVar: Var[Seq[State[T]]] = if (collectIndependent) { Var.collectIndependent(states) } else { Var.collect(states) } - def flip(states: Traversable[State[T]]): State[CC[T]] = { - val notOk = states find { + def flip(states: AnySeq[State[T]]): State[Seq[T]] = { + val notOk = states.find { case Pending | Failed(_) => true case Ok(_) => false } @@ -215,16 +237,16 @@ object Activity { case Some(_) => assert(false) } - val ts = newBuilder() - states foreach { + val ts = new ArrayBuffer[T](states.size) + states.foreach { case Ok(t) => ts += t case _ => assert(false) } - Ok(ts.result) + Ok(ts.toSeq) } - Activity(stateVar map flip) + Activity(stateVar.map(flip)) } /** diff --git a/util-core/src/main/scala/com/twitter/util/Awaitable.scala b/util-core/src/main/scala/com/twitter/util/Awaitable.scala index b0089e0854..4461327d93 100644 --- a/util-core/src/main/scala/com/twitter/util/Awaitable.scala +++ b/util-core/src/main/scala/com/twitter/util/Awaitable.scala @@ -2,7 +2,8 @@ package com.twitter.util import com.twitter.concurrent.Scheduler import java.util.concurrent.atomic.AtomicBoolean -import scala.collection.JavaConverters._ +import scala.annotation.varargs +import scala.jdk.CollectionConverters._ import scala.util.control.{NonFatal => NF} /** @@ -45,7 +46,8 @@ object Awaitable { * * This should only be enabled for threads where blocking is discouraged. * - * @see [[Scheduler.blockingTimeNanos]] and [[Scheduler.blocking]] + * @see [[com.twitter.concurrent.Scheduler.blockingTimeNanos]] and + * [[com.twitter.concurrent.Scheduler.blocking]] */ def enableBlockingTimeTracking(): Unit = TrackElapsedBlocking.set(java.lang.Boolean.TRUE) @@ -59,7 +61,8 @@ object Awaitable { /** * Whether or not this thread should track time spent in blocking operations. * - * @see [[Scheduler.blockingTimeNanos]] and [[Scheduler.blocking]] + * @see [[com.twitter.concurrent.Scheduler.blockingTimeNanos]] and + * [[com.twitter.concurrent.Scheduler.blocking]] */ def getBlockingTimeTracking: Boolean = { TrackElapsedBlocking.get() != null @@ -160,7 +163,7 @@ object Await { /** $all */ @throws(classOf[TimeoutException]) @throws(classOf[InterruptedException]) - @scala.annotation.varargs + @varargs def all(awaitables: Awaitable[_]*): Unit = all(awaitables, Duration.Top) @@ -186,10 +189,10 @@ object Await { all(awaitables.asScala.toSeq, timeout) } -// See http://stackoverflow.com/questions/26643045/java-interoperability-woes-with-scala-generics-and-boxing +// See https://stackoverflow.com/questions/26643045/java-interoperability-woes-with-scala-generics-and-boxing private[util] trait CloseAwaitably0[U <: Unit] extends Awaitable[U] { - private[this] val onClose = new Promise[U] - private[this] val closed = new AtomicBoolean(false) + private[this] val onClose: Promise[U] = new Promise[U] + private[this] val closed: AtomicBoolean = new AtomicBoolean(false) /** * closeAwaitably is intended to be used as a wrapper for @@ -221,6 +224,40 @@ private[util] trait CloseAwaitably0[U <: Unit] extends Awaitable[U] { onClose.isReady } +// See https://stackoverflow.com/questions/26643045/java-interoperability-woes-with-scala-generics-and-boxing +// +// NOTE: the self type must mixin [[Closable]] due to Scala 3 requirement that the self type is closed. +// See https://github.com/lampepfl/dotty/issues/2214 & +// https://github.com/lampepfl/dotty/commit/175499537c87c78d0b926d84b7a9030011e42c00 +private[util] trait CloseOnceAwaitably0[U <: Unit] extends Awaitable[U] { + self: CloseOnce with Closable => + + def ready(timeout: Duration)(implicit permit: Awaitable.CanAwait): this.type = { + closeFuture.ready(timeout) + this + } + + def result(timeout: Duration)(implicit permit: Awaitable.CanAwait): U = + closeFuture.result(timeout).asInstanceOf[U] + + def isReady(implicit permit: Awaitable.CanAwait): Boolean = + closeFuture.isReady +} + +/** + * A mixin to make an [[com.twitter.util.Awaitable]] out + * of a [[com.twitter.util.ClosableOnce]]. + * + * {{{ + * class MyClosable extends ClosableOnce with CloseOnceAwaitably { + * def closeOnce(deadline: Time) = { + * // close the resource + * } + * } + * }}} + */ +trait CloseOnceAwaitably extends CloseOnceAwaitably0[Unit] { self: CloseOnce with Closable => } + /** * A mixin to make an [[com.twitter.util.Awaitable]] out * of a [[com.twitter.util.Closable]]. diff --git a/util-core/src/main/scala/com/twitter/util/Base64Long.scala b/util-core/src/main/scala/com/twitter/util/Base64Long.scala index 59b26e8be3..f2b6c33509 100644 --- a/util-core/src/main/scala/com/twitter/util/Base64Long.scala +++ b/util-core/src/main/scala/com/twitter/util/Base64Long.scala @@ -74,7 +74,11 @@ object Base64Long { * The number is treated as unsigned, so there is never a leading negative sign, and the * representations of negative numbers are larger than positive numbers. */ - def toBase64(builder: StringBuilder, l: Long, alphabet: Int => Char = StandardBase64Alphabet): Unit = { + def toBase64( + builder: StringBuilder, + l: Long, + alphabet: Int => Char = StandardBase64Alphabet + ): Unit = { if (l == 0) { // Special case for zero: Just like in decimal, if the number is // zero, represent it as a single zero digit. @@ -109,9 +113,7 @@ object Base64Long { new PartialFunction[Char, Int] { private[this] val maxChar = chars.max private[this] val reverse = Array.fill[Byte](maxChar.toInt + 1)(-1) - 0.until(AlphabetSize).foreach { i => - reverse(forward(i)) = i.toByte - } + 0.until(AlphabetSize).foreach { i => reverse(forward(i)) = i.toByte } def isDefinedAt(c: Char): Boolean = (c <= maxChar) && (reverse(c) != -1) def apply(c: Char): Int = reverse(c) } @@ -125,16 +127,16 @@ object Base64Long { /** * Decode a sequence of characters representing a base-64-encoded - * Long, as produced by [[toBase64]]. An empty range of characters - * will yield 0L. Leading "zero" digits will not cause overflow. + * Long, as produced by [[com.twitter.util.Base64Long.toBase64(l:Long):String* toBase64(Long)]]. + * An empty range of characters will yield 0L. Leading "zero" digits will not cause overflow. * * @param alphabet defines the mapping between characters in the * string and 6-bit numbers. * - * @throws IllegalArgumentException if any characters in the specified + * @note Throws `IllegalArgumentException` if any characters in the specified * range are not in the specified base-64 alphabet. * - * @throws ArithmeticException if the resulting value overflows a Long. + * @note Throws `ArithmeticException` if the resulting value overflows a Long. */ def fromBase64( chars: CharSequence, diff --git a/util-core/src/main/scala/com/twitter/util/BatchExecutor.scala b/util-core/src/main/scala/com/twitter/util/BatchExecutor.scala index c6e640e661..b89be54dca 100644 --- a/util-core/src/main/scala/com/twitter/util/BatchExecutor.scala +++ b/util-core/src/main/scala/com/twitter/util/BatchExecutor.scala @@ -3,7 +3,8 @@ package com.twitter.util import java.util.concurrent.CancellationException import java.util.logging.Logger -import scala.collection.mutable +import scala.collection.{mutable, Iterable} +import scala.collection.mutable.ArrayBuffer /** * A "batcher" that takes a function `Seq[In] => Future[Seq[Out]]` and @@ -29,15 +30,15 @@ private[util] class BatchExecutor[In, Out]( sizeThreshold: Int, timeThreshold: Duration = Duration.Top, sizePercentile: => Float = 1.0f, - f: Seq[In] => Future[Seq[Out]] + f: Iterable[In] => Future[Seq[Out]] )( - implicit timer: Timer -) extends Function1[In, Future[Out]] { batcher => + implicit timer: Timer) + extends Function1[In, Future[Out]] { batcher => import java.util.logging.Level.WARNING class ScheduledFlush(after: Duration, timer: Timer) { - @volatile var cancelled = false - val task = timer.schedule(after.fromNow) { flush() } + @volatile var cancelled: Boolean = false + val task: TimerTask = timer.schedule(after.fromNow) { flush() } def cancel(): Unit = { cancelled = true @@ -55,14 +56,15 @@ private[util] class BatchExecutor[In, Out]( } } - val log = Logger.getLogger("Future.batched") + val log: Logger = Logger.getLogger("Future.batched") // operations on these are synchronized on `this`. - val buf = new mutable.ArrayBuffer[(In, Promise[Out])](sizeThreshold) + val buf: ArrayBuffer[Tuple2[In, Promise[Out]]] = + new mutable.ArrayBuffer[(In, Promise[Out])](sizeThreshold) var scheduled: Option[ScheduledFlush] = scala.None - var currentBufThreshold = newBufThreshold + var currentBufThreshold: Int = newBufThreshold - def currentBufPercentile = sizePercentile match { + def currentBufPercentile: Float = sizePercentile match { case tooHigh if tooHigh > 1.0f => log.log(WARNING, "value returned for sizePercentile (%f) was > 1.0f, using 1.0", tooHigh) 1.0f @@ -74,7 +76,7 @@ private[util] class BatchExecutor[In, Out]( case p => p } - def newBufThreshold = + def newBufThreshold: Int = math.round(currentBufPercentile * sizeThreshold) match { case tooLow if tooLow < 1 => 1 case size => math.min(size, sizeThreshold) @@ -90,8 +92,7 @@ private[util] class BatchExecutor[In, Out]( flushBatch() else { scheduleFlushIfNecessary() - () => - () + () => () } } @@ -116,7 +117,7 @@ private[util] class BatchExecutor[In, Out]( def flushBatch(): () => Unit = { // this must be executed within a `synchronized` block. val prevBatch = new mutable.ArrayBuffer[(In, Promise[Out])](buf.length) - buf.copyToBuffer(prevBatch) + prevBatch ++= buf buf.clear() scheduled foreach { _.cancel() } @@ -132,8 +133,8 @@ private[util] class BatchExecutor[In, Out]( } } - def executeBatch(batch: Seq[(In, Promise[Out])]): Unit = { - val uncancelled = batch filter { + def executeBatch(batch: Iterable[(In, Promise[Out])]): Unit = { + val uncancelled = batch.filter { case (in, p) => p.isInterrupted match { case Some(_cause) => diff --git a/util-core/src/main/scala/com/twitter/util/Batcher.scala b/util-core/src/main/scala/com/twitter/util/Batcher.scala index 1c1bf9f38e..91bf74dd04 100644 --- a/util-core/src/main/scala/com/twitter/util/Batcher.scala +++ b/util-core/src/main/scala/com/twitter/util/Batcher.scala @@ -1,11 +1,8 @@ package com.twitter.util /** Provides an interface for working with a batched [[com.twitter.util.Future]] */ -class Batcher[In, Out] private[util] ( - executor: BatchExecutor[In, Out] -)( - implicit timer: Timer -) extends (In => Future[Out]) { batcher => +class Batcher[In, Out] private[util] (executor: BatchExecutor[In, Out])(implicit timer: Timer) + extends (In => Future[Out]) { batcher => /** * Enqueues requests for a batched [[com.twitter.util.Future]] diff --git a/util-core/src/main/scala/com/twitter/util/BoundedStack.scala b/util-core/src/main/scala/com/twitter/util/BoundedStack.scala deleted file mode 100644 index 38f8fe7ed6..0000000000 --- a/util-core/src/main/scala/com/twitter/util/BoundedStack.scala +++ /dev/null @@ -1,96 +0,0 @@ -package com.twitter.util - -import scala.reflect.ClassTag - -/** - * A "stack" with a bounded size. If you push a new element on the top - * when the stack is full, the oldest element gets dropped off the bottom. - */ -class BoundedStack[A: ClassTag](val maxSize: Int) extends Seq[A] { - private val array = new Array[A](maxSize) - private var top = 0 - private var count_ = 0 - - def length = count_ - override def size = count_ - - def clear(): Unit = { - top = 0 - count_ = 0 - } - - /** - * Gets the element from the specified index in constant time. - */ - def apply(i: Int): A = { - if (i >= count_) throw new IndexOutOfBoundsException(i.toString) - else array((top + i) % maxSize) - } - - /** - * Pushes an element, possibly forcing out the oldest element in the stack. - */ - def +=(elem: A): Unit = { - top = if (top == 0) maxSize - 1 else top - 1 - array(top) = elem - if (count_ < maxSize) count_ += 1 - } - - /** - * Inserts an element 'i' positions down in the stack. An 'i' value - * of 0 is the same as calling this += elem. This is a O(n) operation - * as elements need to be shifted around. - */ - def insert(i: Int, elem: A): Unit = { - if (i == 0) this += elem - else if (i > count_) throw new IndexOutOfBoundsException(i.toString) - else if (i == count_) { - array((top + i) % maxSize) = elem - count_ += 1 - } else { - val swapped = this(i) - this(i) = elem - insert(i - 1, swapped) - } - } - - /** - * Replaces an element in the stack. - */ - def update(index: Int, elem: A): Unit = { - array((top + index) % maxSize) = elem - } - - /** - * Adds multiple elements, possibly overwriting the oldest elements in - * the stack. If the given iterable contains more elements that this - * stack can hold, then only the last maxSize elements will end up in - * the stack. - */ - def ++=(iter: Iterable[A]): Unit = { - for (elem <- iter) this += elem - } - - /** - * Removes the top element in the stack. - */ - def pop: A = { - if (count_ == 0) throw new NoSuchElementException - else { - val res = array(top) - top = (top + 1) % maxSize - count_ -= 1 - res - } - } - - override def iterator = new Iterator[A] { - var idx = 0 - def hasNext = idx != count_ - def next = { - val res = apply(idx) - idx += 1 - res - } - } -} diff --git a/util-core/src/main/scala/com/twitter/util/Cancellable.scala b/util-core/src/main/scala/com/twitter/util/Cancellable.scala index a1fffb9f12..04d0114ec3 100644 --- a/util-core/src/main/scala/com/twitter/util/Cancellable.scala +++ b/util-core/src/main/scala/com/twitter/util/Cancellable.scala @@ -38,7 +38,7 @@ object Cancellable { class CancellableSink(f: => Unit) extends Cancellable { private[this] val wasCancelled = new AtomicBoolean(false) - def isCancelled = wasCancelled.get + def isCancelled: Boolean = wasCancelled.get def cancel(): Unit = { if (wasCancelled.compareAndSet(false, true)) f } def linkTo(other: Cancellable): Unit = { throw new Exception("linking not supported in CancellableSink") diff --git a/util-core/src/main/scala/com/twitter/util/Closable.scala b/util-core/src/main/scala/com/twitter/util/Closable.scala index 776ad9dc1f..5b75b52302 100644 --- a/util-core/src/main/scala/com/twitter/util/Closable.scala +++ b/util-core/src/main/scala/com/twitter/util/Closable.scala @@ -1,10 +1,15 @@ package com.twitter.util -import java.lang.ref.{PhantomReference, Reference, ReferenceQueue} +import java.lang.ref.PhantomReference +import java.lang.ref.Reference +import java.lang.ref.ReferenceQueue import java.{util => ju} +import java.util.concurrent.atomic.AtomicBoolean import java.util.concurrent.atomic.AtomicReference -import java.util.logging.{Level, Logger} +import java.util.logging.Level +import java.util.logging.Logger import scala.annotation.tailrec +import scala.annotation.varargs import scala.util.control.{NonFatal => NF} /** @@ -45,6 +50,15 @@ abstract class AbstractClosable extends Closable */ object Closable { private[this] val logger = Logger.getLogger("") + private[this] val collectorThreadEnabled = new AtomicBoolean(true) + + /** Stop CollectClosables thread. */ + def stopCollectClosablesThread(): Unit = { + if (collectorThreadEnabled.compareAndSet(true, false)) + Try(collectorThread.interrupt).onFailure(fatal => + logger + .log(Level.SEVERE, "Current thread cannot interrupt CollectClosables thread", fatal))() + } /** Provide Java access to the [[com.twitter.util.Closable]] mixin. */ def close(o: AnyRef): Future[Unit] = o match { @@ -72,23 +86,21 @@ object Closable { * Concurrent composition: creates a new closable which, when * closed, closes all of the underlying resources simultaneously. */ - def all(closables: Closable*): Closable = new Closable { + @varargs def all(closables: Closable*): Closable = new Closable { def close(deadline: Time): Future[Unit] = { - val fs = closables.map { closable => - safeClose(closable, deadline) - } - val iter = fs.iterator + val iter = closables.iterator @tailrec - def checkNext(): Future[Unit] = { - if (!iter.hasNext) Future.Done + def checkNext(async: List[Future[Unit]]): List[Future[Unit]] = { + if (!iter.hasNext) async else { - iter.next().poll match { - case Some(Return(_)) => checkNext() - case _ => Future.join(fs) + safeClose(iter.next(), deadline) match { + case Future.Done => checkNext(async) + case f => checkNext(f :: async) } } } - checkNext() + + Future.join(checkNext(Nil)) } } @@ -98,20 +110,58 @@ object Closable { * resource ''n+1'' is not closed until resource ''n'' is. * * @return the first failed [[Future]] should any of the `Closables` - * result in a failed [[Future]]. + * result in a failed [[Future]]. Any subsequent failures are + * included as suppressed exceptions. * * @note as with all `Closables`, the `deadline` passed to `close` * is advisory. */ + def sequence(a: Closable, b: Closable): Closable = new Closable { + def close(deadline: Time): Future[Unit] = { + val first = safeClose(a, deadline) + + def onFirstComplete(t: Try[Unit]): Future[Unit] = { + val second = safeClose(b, deadline) + if (t.isThrow) { + if (second eq Future.Done) first + else { + second.transform { + case Return(_) => first + case Throw(secondThrow) => + t.throwable.addSuppressed(secondThrow) + first + } + } + } else second + } + + if (first eq Future.Done) onFirstComplete(Try.Unit) + else first.transform(onFirstComplete) + } + } + + /** + * Sequential composition: create a new Closable which, when + * closed, closes all of the underlying ones in sequence: that is, + * resource ''n+1'' is not closed until resource ''n'' is. + * + * @return the first failed [[Future]] should any of the `Closables` + * result in a failed [[Future]]. Any subsequent failures are + * included as suppressed exceptions. + * + * @note as with all `Closables`, the `deadline` passed to `close` + * is advisory. + */ + @varargs def sequence(closables: Closable*): Closable = new Closable { private final def closeSeq( deadline: Time, closables: Seq[Closable], - firstFailure: Option[Future[Unit]] + firstFailure: Option[Throw[Unit]] ): Future[Unit] = closables match { case Seq() => firstFailure match { - case Some(f) => f + case Some(t) => Future.const(t) case None => Future.Done } case Seq(hd, tl @ _*) => @@ -120,8 +170,10 @@ object Closable { case Return(_) => firstFailure case t @ Throw(_) => firstFailure match { - case Some(_) => firstFailure - case None => Some(Future.const(t)) + case Some(firstThrow) => + firstThrow.throwable.addSuppressed(t.throwable) + firstFailure + case None => Some(t) } } closeSeq(deadline, tl, failure) @@ -129,10 +181,8 @@ object Closable { val firstClose = safeClose(hd, deadline) // eagerness is needed by `Var`, see RB_ID=557464 - firstClose.poll match { - case Some(t) => onTry(t) - case None => firstClose.transform(onTry) - } + if (firstClose eq Future.Done) onTry(Try.Unit) + else firstClose.transform(onTry) } def close(deadline: Time): Future[Unit] = @@ -160,7 +210,7 @@ object Closable { private val collectorThread = new Thread("CollectClosables") { override def run(): Unit = { - while (true) { + while (collectorThreadEnabled.get()) { try { val ref = refq.remove() val closable = refs.synchronized(refs.remove(ref)) diff --git a/util-core/src/main/scala/com/twitter/util/ClosableOnce.scala b/util-core/src/main/scala/com/twitter/util/ClosableOnce.scala index 8ea5b0e216..c40490dfa6 100644 --- a/util-core/src/main/scala/com/twitter/util/ClosableOnce.scala +++ b/util-core/src/main/scala/com/twitter/util/ClosableOnce.scala @@ -1,49 +1,16 @@ package com.twitter.util -import scala.util.control.NonFatal - /** * A mixin trait to describe resources that are an idempotent [[Closable]]. * * The first call to `close(Time)` triggers closing the resource with the * provided deadline while subsequent calls will yield the same `Future` * as the first invocation irrespective of the deadline provided. + * + * @see [[CloseOnce]] if you are mixing in or extending an existing [[Closable]] + * to prevent any "inherits from conflicting members" compile-time issues. */ -trait ClosableOnce extends Closable { - - // Our intrinsic lock for mutating the `closed` field - private[this] val closePromise = Promise[Unit]() - private[this] var closed = false - - private[this] def once(): Boolean = closePromise.synchronized { - if (closed) { - false - } else { - closed = true - true - } - } - - /** - * Close the resource with the given deadline exactly once. This deadline is - * advisory, giving the callee some leeway, for example to drain clients or - * finish up other tasks. - * - * @note if this method throws a synchronous exception, that exception will - * be wrapped in a failed future. - */ - protected def closeOnce(deadline: Time): Future[Unit] - - final def close(deadline: Time): Future[Unit] = { - if (once()) { - val closeF = - try closeOnce(deadline) - catch { case NonFatal(ex) => Future.exception(ex) } - closePromise.become(closeF) - } - closePromise - } -} +trait ClosableOnce extends Closable with CloseOnce object ClosableOnce { diff --git a/util-core/src/main/scala/com/twitter/util/CloseOnce.scala b/util-core/src/main/scala/com/twitter/util/CloseOnce.scala new file mode 100644 index 0000000000..eb4a274e4c --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/CloseOnce.scala @@ -0,0 +1,78 @@ +package com.twitter.util + +import scala.util.control.NonFatal + +/** + * A mixin trait to describe resources that are an idempotent [[Closable]]. + * + * The first call to `close(Time)` triggers closing the resource with the + * provided deadline while subsequent calls will yield the same `Future` + * as the first invocation irrespective of the deadline provided. + * + * @see [[ClosableOnce]] if you are not mixing in or extending an existing [[Closable]] + * @see [[ClosableOnce.of]] for creating a proxy to a [[Closable]] + * that has already been instantiated. + * + * @note this mixin extends [[Closable]] as a workaround for extending [[AbstractClosableOnce]] + * in Java. The mixin's final override is not recognized when compiled with Scala 3. See CSL-11187 + * or https://github.com/lampepfl/dotty/issues/13104. + */ +trait CloseOnce extends Closable { self: Closable => + // Our intrinsic lock for mutating the `closed` field + private[this] val closePromise: Promise[Unit] = Promise[Unit]() + @volatile private[this] var closed: Boolean = false + + // optimized to limit synchronized locking scope + private[this] def firstCloseInvoked(): Boolean = closePromise.synchronized { + if (closed) { + false + } else { + closed = true + true + } + } + + /** + * Close the resource with the given deadline exactly once. This deadline is + * advisory, giving the callee some leeway, for example to drain clients or + * finish up other tasks. + * + * @note if this method throws a synchronous exception, that exception will + * be wrapped in a failed future. If a fatal error is thrown synchronously, + * the error will be wrapped in a failed Future AND thrown directly. + */ + protected def closeOnce(deadline: Time): Future[Unit] + + /** + * The [[Future]] satisfied upon completion of [[close]]. + * + * @note we do not expose direct access to the underlying [[Promise closePromise]], because + * the [[Promise]] state is mutable - we only want mutation to occur within this + * [[CloseOnce]] trait. + */ + protected final def closeFuture: Future[Unit] = closePromise + + /** + * Signals whether or not this [[Closable]] has been closed. + * + * @return return true if [[close]] has been initiated, false otherwise + */ + final def isClosed: Boolean = closed + + override final def close(deadline: Time): Future[Unit] = { + // only call `closeOnce()` and assign `closePromise` if this is the first `close()` invocation + if (firstCloseInvoked()) { + try { + closePromise.become(closeOnce(deadline)) + } catch { + case NonFatal(ex) => + closePromise.setException(ex) + case t: Throwable => + closePromise.setException(t) + throw t + } + } + + closePromise + } +} diff --git a/util-core/src/main/scala/com/twitter/util/Codec.scala b/util-core/src/main/scala/com/twitter/util/Codec.scala index 48d5bac5ac..71ee2bc205 100644 --- a/util-core/src/main/scala/com/twitter/util/Codec.scala +++ b/util-core/src/main/scala/com/twitter/util/Codec.scala @@ -25,7 +25,7 @@ trait Encoder[T, S] extends (T => S) { */ def encode(t: T): S - def apply(t: T) = encode(t) + def apply(t: T): S = encode(t) } trait DecoderCompanion { @@ -51,7 +51,7 @@ trait Decoder[T, S] { */ def decode(s: S): T - def invert(s: S) = decode(s) + def invert(s: S): T = decode(s) } object Codec extends EncoderCompanion with DecoderCompanion { diff --git a/util-core/src/main/scala/com/twitter/util/Config.scala b/util-core/src/main/scala/com/twitter/util/Config.scala deleted file mode 100644 index 911cc548bc..0000000000 --- a/util-core/src/main/scala/com/twitter/util/Config.scala +++ /dev/null @@ -1,187 +0,0 @@ -/* - * Copyright 2011 Twitter, Inc. - * - * 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 com.twitter.util - -import java.io.Serializable -import scala.collection.mutable -import scala.language.implicitConversions - -/** - * You can import Config._ if you want the auto-conversions in a class - * that does not inherit from trait Config. - */ -@deprecated("use a Plain Old Scala Object", "") -object Config { - sealed trait Required[+A] extends Serializable { - def value: A - def isSpecified: Boolean - def isEmpty: Boolean = !isSpecified - } - - object Specified { - def unapply[A](specified: Specified[A]): Some[A] = Some(specified.value) - } - - class Specified[+A](f: => A) extends Required[A] { - lazy val value: A = f - def isSpecified: Boolean = true - } - - case object Unspecified extends Required[scala.Nothing] { - def value: scala.Nothing = throw new NoSuchElementException - def isSpecified: Boolean = false - } - - class RequiredValuesMissing(names: Seq[String]) extends Exception(names.mkString(",")) - - /** - * Config classes that don't produce anything via an apply() method - * can extends Config.Nothing, and still get access to the Required - * classes and methods. - */ - class Nothing extends Config[scala.Nothing] { - def apply(): scala.Nothing = throw new UnsupportedOperationException - } - - implicit def toSpecified[A](value: => A): Specified[A] = new Specified(value) - implicit def toSpecifiedOption[A](value: => A): Specified[Some[A]] = - new Specified(Some(value)) - implicit def fromRequired[A](req: Required[A]): A = req.value - implicit def intoOption[A](item: A): Option[A] = Some(item) - implicit def intoList[A](item: A): List[A] = List(item) -} - -/** - * A config object contains vars for configuring an object of type T, and an apply() method - * which turns this config into an object of type T. - * - * The trait also contains a few useful implicits. - * - * You can define fields that are required but that don't have a default value with: - * class BotConfig extends Config[Bot] { - * var botName = required[String] - * var botPort = required[Int] - * .... - * } - * - * Optional fields can be defined with: - * var something = optional[Duration] - * - * Fields that are dependent on other fields and have a default value computed - * from an expression should be marked as computed: - * - * var level = required[Int] - * var nextLevel = computed { level + 1 } - */ -@deprecated("use a Plain Old Scala Object", "") -trait Config[T] extends (() => T) { - import Config.{Required, Specified, Unspecified, RequiredValuesMissing} - - /** - * a lazy singleton of the configured class. Allows one to pass - * the Config around to multiple classes without producing multiple - * configured instances - */ - lazy val memoized: T = apply() - - def required[A]: Required[A] = Unspecified - def required[A](default: => A): Required[A] = new Specified(default) - def optional[A]: Required[Option[A]] = new Specified(None) - def optional[A](default: => A): Required[Option[A]] = new Specified(Some(default)) - - /** - * The same as specifying required[A] with a default value, but the intent to the - * use is slightly different. - */ - def computed[A](f: => A): Required[A] = new Specified(f) - - /** a non-lazy way to create a Specified that can be called from java */ - def specified[A](value: A): Specified[A] = new Specified(value) - - implicit def toSpecified[A](value: => A): Specified[A] = new Specified(value) - implicit def toSpecifiedOption[A](value: => A): Specified[Some[A]] = - new Specified(Some(value)) - implicit def fromRequired[A](req: Required[A]): A = req.value - implicit def intoOption[A](item: A): Option[A] = Some(item) - implicit def intoList[A](item: A): List[A] = List(item) - - /** - * Looks for any fields that are Required but currently Unspecified, in - * this object and in any child Config objects, and returns the names - * of those missing values. For missing values in child Config objects, - * the name is the dot-separated path down to that value. - * - * @note as of Scala 2.12 the returned values are not always correct. - * See `ConfigTest` for an example. - */ - def missingValues: Seq[String] = { - val alreadyVisited = mutable.Set[AnyRef]() - val interestingReturnTypes = Seq(classOf[Required[_]], classOf[Config[_]], classOf[Option[_]]) - val buf = new mutable.ListBuffer[String] - val mirror = scala.reflect.runtime.universe.runtimeMirror(getClass.getClassLoader) - def collect(prefix: String, config: Config[_]): Unit = { - if (!alreadyVisited.contains(config)) { - alreadyVisited += config - val nullaryMethods = config.getClass.getMethods.toSeq filter { _.getParameterTypes.isEmpty } - val syntheticTermNames: Set[String] = { - val configSymbol = mirror.reflect(config).symbol - val configMembers = configSymbol.toType.members - configMembers.collect { - case symbol if symbol.isTerm && symbol.isSynthetic => - symbol.name.decodedName.toString - }(scala.collection.breakOut) - } - for (method <- nullaryMethods) { - val name = method.getName - val rt = method.getReturnType - if (name != "required" && - name != "optional" && - !syntheticTermNames.contains(name) && // no loops, no compiler generated methods! - interestingReturnTypes.exists(_.isAssignableFrom(rt))) { - method.invoke(config) match { - case Unspecified => - buf += (prefix + name) - case specified: Specified[_] => - specified match { - case Specified(sub: Config[_]) => - collect(prefix + name + ".", sub) - case Specified(Some(sub: Config[_])) => - collect(prefix + name + ".", sub) - case _ => - } - case sub: Config[_] => - collect(prefix + name + ".", sub) - case Some(sub: Config[_]) => - collect(prefix + name + ".", sub) - case _ => - } - } - } - } - } - collect("", this) - buf.toList - } - - /** - * @throws RequiredValuesMissing if there are any missing values. - */ - def validate(): Unit = { - val missing = missingValues - if (missing.nonEmpty) throw new RequiredValuesMissing(missing) - } -} diff --git a/util-core/src/main/scala/com/twitter/util/ConstFuture.scala b/util-core/src/main/scala/com/twitter/util/ConstFuture.scala index a43fa82b16..3659208df7 100644 --- a/util-core/src/main/scala/com/twitter/util/ConstFuture.scala +++ b/util-core/src/main/scala/com/twitter/util/ConstFuture.scala @@ -1,6 +1,5 @@ package com.twitter.util -import com.twitter.concurrent.Scheduler import scala.runtime.NonLocalReturnControl /** @@ -17,31 +16,55 @@ class ConstFuture[A](result: Try[A]) extends Future[A] { // https://twitter.github.io/util/guide/util-cookbook/futures.html#future-recursion // for details. The second is that this keeps the execution order consistent // with `Promise`. - def respond(k: Try[A] => Unit): Future[A] = { + def respond(f: Try[A] => Unit): Future[A] = { val saved = Local.save() - Scheduler.submit(new Runnable { - def run(): Unit = { - val current = Local.save() - Local.restore(saved) - try k(result) - catch Monitor.catcher - finally Local.restore(current) - } + saved.fiber.submitTask(() => { + val current = Local.save() + if (current ne saved) Local.restore(saved) + + val tracker = saved.resourceTracker + val run = + if (tracker eq None) f + else ResourceTracker.wrapAndMeasureUsage(f, tracker.get) + + try run(result) + catch Monitor.catcher + finally Local.restore(current) }) this } + override def proxyTo[B >: A](other: Promise[B]): Unit = { + // avoid an extra call to `isDefined` as `update` checks + other.update(result) + } + def raise(interrupt: Throwable): Unit = () + // NOTE: since `f` is not run by the scheduler, the used cpu time will be tracked by + // the continuation that executes this transformation + protected def transformTry[B](f: Try[A] => Try[B]): Future[B] = { + try Future.const(f(result)) + catch { + case scala.util.control.NonFatal(e) => Future.const(Throw(e)) + } + } + def transform[B](f: Try[A] => Future[B]): Future[B] = { val p = new Promise[B] // see the note on `respond` for an explanation of why `Scheduler` is used. val saved = Local.save() - Scheduler.submit(new Runnable { - def run(): Unit = { - val current = Local.save() - Local.restore(saved) - val computed = try f(result) + saved.fiber.submitTask(() => { + val current = Local.save() + if (current ne saved) Local.restore(saved) + + val tracker = saved.resourceTracker + val run = + if (tracker eq None) f + else ResourceTracker.wrapAndMeasureUsage(f, tracker.get) + + val computed = + try run(result) catch { case e: NonLocalReturnControl[_] => Future.exception(new FutureNonLocalReturnControl(e)) case scala.util.control.NonFatal(e) => Future.exception(e) @@ -49,8 +72,8 @@ class ConstFuture[A](result: Try[A]) extends Future[A] { Monitor.handle(t) throw t } finally Local.restore(current) - p.become(computed) - } + + p.become(computed) }) p } @@ -59,8 +82,6 @@ class ConstFuture[A](result: Try[A]) extends Future[A] { override def isDefined: Boolean = true - override def isDone(implicit ev: this.type <:< Future[Unit]): Boolean = result.isReturn - override def toString: String = s"ConstFuture($result)" // Awaitable diff --git a/util-core/src/main/scala/com/twitter/util/CountDownLatch.scala b/util-core/src/main/scala/com/twitter/util/CountDownLatch.scala deleted file mode 100644 index 6e885ee5af..0000000000 --- a/util-core/src/main/scala/com/twitter/util/CountDownLatch.scala +++ /dev/null @@ -1,15 +0,0 @@ -package com.twitter.util - -import java.util.concurrent.TimeUnit - -class CountDownLatch(val initialCount: Int) { - val underlying = new java.util.concurrent.CountDownLatch(initialCount) - def count = underlying.getCount - def isZero = count == 0 - def countDown() = underlying.countDown() - def await() = underlying.await() - def await(timeout: Duration) = underlying.await(timeout.inMillis, TimeUnit.MILLISECONDS) - def within(timeout: Duration) = await(timeout) || { - throw new TimeoutException(timeout.toString) - } -} diff --git a/util-core/src/main/scala/com/twitter/util/Credentials.scala b/util-core/src/main/scala/com/twitter/util/Credentials.scala deleted file mode 100644 index c535164516..0000000000 --- a/util-core/src/main/scala/com/twitter/util/Credentials.scala +++ /dev/null @@ -1,69 +0,0 @@ -/* - * Copyright 2011 Twitter, Inc. - * - * 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 com.twitter.util - -import java.io.{File, IOException} - -import scala.collection.JavaConverters._ -import scala.io.Source -import scala.util.parsing.combinator._ - -/** - * Simple helper to read authentication credentials from a text file. - * - * The file's format is assumed to be trivialized yaml, containing lines of the form ``key: value``. - * Keys can be any word character or '-' and values can be any character except new lines. - */ -object Credentials { - object parser extends RegexParsers { - - override val whiteSpace = "(?:\\s+|#.*\\r?\\n)+".r - - private[this] val key = "[\\w-]+".r - private[this] val value = ".+".r - - def auth = key ~ ":" ~ value ^^ { case k ~ ":" ~ v => (k, v) } - def content: Parser[Map[String, String]] = rep(auth) ^^ { auths => - Map(auths: _*) - } - - def apply(in: String): Map[String, String] = { - parseAll(content, in) match { - case Success(result, _) => result - case x: Failure => throw new IOException(x.toString) - case x: Error => throw new IOException(x.toString) - } - } - } - - def apply(file: File): Map[String, String] = parser(Source.fromFile(file).mkString) - - def apply(data: String): Map[String, String] = parser(data) - - def byName(name: String): Map[String, String] = { - apply(new File(System.getenv().asScala.getOrElse("KEY_FOLDER", "/etc/keys"), name)) - } -} - -/** - * Java interface to Credentials. - */ -class Credentials { - def read(data: String): java.util.Map[String, String] = Credentials(data).asJava - def read(file: File): java.util.Map[String, String] = Credentials(file).asJava - def byName(name: String): java.util.Map[String, String] = Credentials.byName(name).asJava -} diff --git a/util-core/src/main/scala/com/twitter/util/Diff.scala b/util-core/src/main/scala/com/twitter/util/Diff.scala index 8090a40a81..1a58fdc29e 100644 --- a/util-core/src/main/scala/com/twitter/util/Diff.scala +++ b/util-core/src/main/scala/com/twitter/util/Diff.scala @@ -71,7 +71,7 @@ object Diffable { def map[U](f: T => U): SetDiff[U] = SetDiff(add.map(f), remove.map(f)) - override def toString = s"Diff(+$add, -$remove)" + override def toString: String = s"Diff(+$add, -$remove)" } private case class SeqDiff[T](limit: Int, insert: Map[Int, T]) extends Diff[Seq, T] { @@ -86,7 +86,7 @@ object Diffable { } def map[U](f: T => U): SeqDiff[U] = - SeqDiff(limit, insert.mapValues(f)) + SeqDiff(limit, insert.map { case (k, v) => (k, f(v)) }) } implicit val ofSet: Diffable[Set] = new Diffable[Set] { @@ -98,14 +98,17 @@ object Diffable { def diff[T](left: Seq[T], right: Seq[T]): SeqDiff[T] = if (left.length < right.length) { val SeqDiff(_, insert) = diff(left, right.take(left.length)) - SeqDiff(right.length, insert ++ ((left.length until right.length) map { i => - (i -> right(i)) - })) + SeqDiff( + right.length, + insert ++ ((left.length until right.length) map { i => + (i -> right(i)) + })) } else if (left.length > right.length) { diff(left.take(right.length), right) } else { - val insert = for (((x, y), i) <- left.zip(right).zipWithIndex if x != y) - yield (i -> y) + val insert = + for (((x, y), i) <- left.zip(right).zipWithIndex if x != y) + yield (i -> y) SeqDiff(left.length, insert.toMap) } diff --git a/util-core/src/main/scala/com/twitter/util/Disposable.scala b/util-core/src/main/scala/com/twitter/util/Disposable.scala index b4cc8ef7e3..78d29bae7f 100644 --- a/util-core/src/main/scala/com/twitter/util/Disposable.scala +++ b/util-core/src/main/scala/com/twitter/util/Disposable.scala @@ -66,7 +66,7 @@ trait Disposable[+T] { } object Disposable { - def const[T](t: T) = new Disposable[T] { + def const[T](t: T): Disposable[T] = new Disposable[T] { def get = t def dispose(deadline: Time) = Future.value(()) } @@ -97,13 +97,14 @@ trait Managed[+T] { selfT => def make() = new Disposable[U] { val t = selfT.make() - val u = try { - f(t.get).make() - } catch { - case e: Exception => - t.dispose() - throw e - } + val u = + try { + f(t.get).make() + } catch { + case e: Exception => + t.dispose() + throw e + } def get = u.get @@ -111,7 +112,7 @@ trait Managed[+T] { selfT => u.dispose(deadline) transform { case Return(_) => t.dispose(deadline) case Throw(outer) => - t.dispose transform { + t.dispose() transform { case Throw(inner) => Future.exception(new DoubleTrouble(outer, inner)) case Return(_) => Future.exception(outer) } @@ -120,9 +121,7 @@ trait Managed[+T] { selfT => } } - def map[U](f: T => U): Managed[U] = flatMap { t => - Managed.const(f(t)) - } + def map[U](f: T => U): Managed[U] = flatMap { t => Managed.const(f(t)) } /** * Builds a resource. @@ -131,13 +130,13 @@ trait Managed[+T] { selfT => } object Managed { - def singleton[T](t: Disposable[T]) = new Managed[T] { def make() = t } - def const[T](t: T) = singleton(Disposable.const(t)) + def singleton[T](t: Disposable[T]): Managed[T] = new Managed[T] { def make() = t } + def const[T](t: T): Managed[T] = singleton(Disposable.const(t)) } class DoubleTrouble(cause1: Throwable, cause2: Throwable) extends Exception { - override def getStackTrace = cause1.getStackTrace - override def getMessage = + override def getStackTrace: Array[StackTraceElement] = cause1.getStackTrace + override def getMessage: String = "Double failure while disposing composite resource: %s \n %s".format( cause1.getMessage, cause2.getMessage diff --git a/util-core/src/main/scala/com/twitter/util/Duration.scala b/util-core/src/main/scala/com/twitter/util/Duration.scala index 188ddca04c..c52403e9ba 100644 --- a/util-core/src/main/scala/com/twitter/util/Duration.scala +++ b/util-core/src/main/scala/com/twitter/util/Duration.scala @@ -4,39 +4,48 @@ import java.io.Serializable import java.util.concurrent.TimeUnit object Duration extends TimeLikeOps[Duration] { - - def fromNanoseconds(nanoseconds: Long): Duration = new Duration(nanoseconds) + def fromNanoseconds(nanoseconds: Long): Duration = + if (nanoseconds == 0L) Zero + else new Duration(nanoseconds) // This is needed for Java compatibility. override def fromFractionalSeconds(seconds: Double): Duration = super.fromFractionalSeconds(seconds) override def fromSeconds(seconds: Int): Duration = super.fromSeconds(seconds) + override def fromMinutes(minutes: Int): Duration = super.fromMinutes(minutes) override def fromMilliseconds(millis: Long): Duration = super.fromMilliseconds(millis) override def fromMicroseconds(micros: Long): Duration = super.fromMicroseconds(micros) + override def fromHours(hours: Int): Duration = super.fromHours(hours) + override def fromDays(days: Int): Duration = super.fromDays(days) - val NanosPerMicrosecond = 1000L - val NanosPerMillisecond = NanosPerMicrosecond * 1000L - val NanosPerSecond = NanosPerMillisecond * 1000L - val NanosPerMinute = NanosPerSecond * 60 - val NanosPerHour = NanosPerMinute * 60 - val NanosPerDay = NanosPerHour * 24 + val NanosPerMicrosecond: Long = 1000L + val NanosPerMillisecond: Long = NanosPerMicrosecond * 1000L + val NanosPerSecond: Long = NanosPerMillisecond * 1000L + val NanosPerMinute: Long = NanosPerSecond * 60 + val NanosPerHour: Long = NanosPerMinute * 60 + val NanosPerDay: Long = NanosPerHour * 24 /** - * Create a duration from a [[java.util.concurrent.TimeUnit]]. - * Synonym for `apply`. + * Create a duration from a `java.util.concurrent.TimeUnit`. + * Synonym for [[apply]]. */ def fromTimeUnit(value: Long, unit: TimeUnit): Duration = apply(value, unit) /** - * Create a duration from a [[java.util.concurrent.TimeUnit]]. + * Create a duration from a `java.util.concurrent.TimeUnit`. */ def apply(value: Long, unit: TimeUnit): Duration = { val ns = TimeUnit.NANOSECONDS.convert(value, unit) fromNanoseconds(ns) } + /** + * Create a duration from a [[java.time.Duration]]. + */ + def fromJava(value: java.time.Duration): Duration = fromNanoseconds(value.toNanos) + // This is needed for Java compatibility. - override val Zero: Duration = fromNanoseconds(0) + override val Zero: Duration = new Duration(0) /** * Duration `Top` is greater than any other duration, except for @@ -218,7 +227,7 @@ object Duration extends TimeLikeOps[Duration] { * * The parser is case-insensitive. * - * @throws RuntimeException if the string cannot be parsed. + * @note Throws `RuntimeException` if the string cannot be parsed. */ def parse(s: String): Duration = { val ss = s.toLowerCase @@ -287,7 +296,7 @@ private[util] object DurationBox { * using the time conversions: * * {{{ - * import com.twitter.conversions.time._ + * import com.twitter.conversions.DurationOps._ * * 3.days+4.nanoseconds * }}} @@ -301,10 +310,11 @@ private[util] object DurationBox { * their arithmetic follows. This is useful for representing durations * that are truly infinite; for example the absence of a timeout. */ -sealed class Duration private[util] (protected val nanos: Long) extends { - protected val ops = Duration -} with TimeLike[Duration] with Serializable { - import ops._ +sealed class Duration private[util] (protected val nanos: Long) + extends TimeLike[Duration] + with Serializable { + override def ops: TimeLikeOps[Duration] = Duration + import Duration._ def inNanoseconds: Long = nanos @@ -317,12 +327,18 @@ sealed class Duration private[util] (protected val nanos: Long) extends { def inUnit(unit: TimeUnit): Long = unit.convert(inNanoseconds, TimeUnit.NANOSECONDS) + /** + * Converts this value to a [[java.time.Duration]]. + */ + def asJava: java.time.Duration = + java.time.Duration.ofNanos(inNanoseconds) + /** * toString produces a representation that * * - loses no information * - is easy to read - * - can be read back in if com.twitter.conversions.time._ is imported + * - can be read back in if com.twitter.conversions.DurationOps._ is imported * * An example: * @@ -421,8 +437,11 @@ sealed class Duration private[util] (protected val nanos: Long) extends { def %(x: Duration): Duration = x match { case Undefined => Undefined case ns if ns.isZero => Undefined - case ns if ns.isFinite => fromNanoseconds(nanos % ns.inNanoseconds) - case Top | Bottom => this + // Analogous to `case Top | Bottom =>`. + // We pattern match this way so that the compiler + // recognizes a fully exhaustive pattern match. + case ns if !ns.isFinite => this + case ns => fromNanoseconds(nanos % ns.inNanoseconds) } /** diff --git a/util-core/src/main/scala/com/twitter/util/Event.scala b/util-core/src/main/scala/com/twitter/util/Event.scala index 5683afc558..39b9ab4591 100644 --- a/util-core/src/main/scala/com/twitter/util/Event.scala +++ b/util-core/src/main/scala/com/twitter/util/Event.scala @@ -3,19 +3,64 @@ package com.twitter.util import java.lang.ref.WeakReference import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} import scala.annotation.tailrec -import scala.collection.generic.CanBuild +import scala.collection.compat._ import scala.collection.immutable.Queue import scala.language.higherKinds /** - * Events are instantaneous values, defined only at particular - * instants in time (cf. [[com.twitter.util.Var Vars]], which are - * defined at all times). It is possible to view Events as the - * discrete counterpart to [[com.twitter.util.Var Var]]'s continuous - * nature. + * Events are emitters of values. They can be processed as if they were + * streams. Here's how to create one that emits Ints. * - * Events are observed by registering [[com.twitter.util.Witness Witnesses]] - * to which the Event's values are notified. + * {{{ + * val ints = Event[Int]() + * }}} + * + * To listen to an Event stream for values, register a handler. + * + * {{{ + * val you = ints.respond { n => println(s"\$n for you") } + * }}} + * + * Handlers in Event-speak are called [[Witness Witnesses]]. + * And, respond is equivalent to `register(Witness { ... })`, so we can also + * write: + * + * {{{ + * val me = ints.register(Witness { n => println(s"\$n for me") }) + * }}} + * + * Registration is time-sensitive. Witnesses registered to an Event stream + * can't observe past values. In this way, Events are instantaneous. When the + * stream emits a value, it fans out to every Witness; if there are no + * Witnesses to that communication, whatever value it carried is lost. + * + * Events emit values, okay, but who decides what is emitted? The implementor + * of an Event decides! + * + * But there's another way: you can create an Event that is also a Witness. + * Remember, Witnesses are simply handlers that can receive a value. Their type + * is `T => Unit`. So an Event that you can send values to is also a Witness. + * + * The `ints` Event we created above is one of these. Watch: + * + * {{{ + * scala> ints.notify(1) + * 1 for you + * 1 for me + * }}} + * + * We can also stop listening. This is important for resource cleanup. + * + * {{{ + * scala> me.close() + * scala> ints.notify(2) + * 2 for you + * }}} + * + * In relation to [[com.twitter.util.Var Vars]], Events are instantaneous + * values, defined only at particular instants in time; whereas Vars are + * defined at all times. It is possible to view Events as the discrete + * counterpart to Var's continuous nature. * * Note: There is a Java-friendly API for this trait: [[com.twitter.util.AbstractEvent]]. */ @@ -41,9 +86,7 @@ trait Event[+T] { self => */ def collect[U](f: PartialFunction[T, U]): Event[U] = new Event[U] { def register(s: Witness[U]): Closable = - self.respond { t => - f.runWith(s.notify)(t) - } + self.respond { t => f.runWith(s.notify)(t) } } /** @@ -108,9 +151,7 @@ trait Event[+T] { self => def mergeMap[U](f: T => Event[U]): Event[U] = new Event[U] { def register(s: Witness[U]): Closable = { @volatile var inners = Nil: List[Closable] - val outer = self.respond { el => - inners.synchronized { inners ::= f(el).register(s) } - } + val outer = self.respond { el => inners.synchronized { inners ::= f(el).register(s) } } Closable.make { deadline => outer.close(deadline) before { @@ -125,12 +166,8 @@ trait Event[+T] { self => */ def select[U](other: Event[U]): Event[Either[T, U]] = new Event[Either[T, U]] { def register(s: Witness[Either[T, U]]): Closable = Closable.all( - self.register(s.comap { t => - Left(t) - }), - other.register(s.comap { u => - Right(u) - }) + self.register(s.comap { t => Left(t) }), + other.register(s.comap { u => Right(u) }) ) } @@ -159,6 +196,8 @@ trait Event[+T] { self => if (rest.isEmpty) state = None else state = Some(Right(Queue(rest: _*))) s.notify((t, u)) + case Some(Right(_)) => + throw new IllegalStateException("unexpected empty queue") } } } @@ -174,6 +213,8 @@ trait Event[+T] { self => if (rest.isEmpty) state = None else state = Some(Left(Queue(rest: _*))) s.notify((t, u)) + case Some(Left(_)) => + throw new IllegalStateException("unexpected empty queue") } } } @@ -259,13 +300,13 @@ trait Event[+T] { self => } /** - * Progressively build a collection of events using the passed-in + * Progressively build any collection of events using the passed-in * builder. A value containing the current version of the collection * is notified for each incoming event. */ - def build[U >: T, That](implicit cbf: CanBuild[U, That]) = new Event[That] { + def buildAny[That](implicit factory: Factory[T, That]): Event[That] = new Event[That] { def register(s: Witness[That]): Closable = { - val b = cbf() + val b = factory.newBuilder self.respond { t => b += t s.notify(b.result()) @@ -273,6 +314,14 @@ trait Event[+T] { self => } } + /** + * Progressively build a Seq of events using the passed-in + * builder. A value containing the current version of the collection + * is notified for each incoming event. + */ + def build[U >: T, That <: Seq[U]](implicit factory: Factory[U, That]): Event[That] = + buildAny(factory) + /** * A Future which is satisfied by the first value observed. */ @@ -339,9 +388,7 @@ trait Event[+T] { self => * Builds a new Event by keeping only the Events where * the previous and current values are not `==` to each other. */ - def dedup: Event[T] = dedupWith { (a, b) => - a == b - } + def dedup: Event[T] = dedupWith { (a, b) => a == b } } /** @@ -430,11 +477,14 @@ abstract class AbstractWitness[T] extends Witness[T] /** * Note: There is Java-friendly API for this object: [[com.twitter.util.Witnesses]]. + * + * @define VarApplyScaladocLink + * [[com.twitter.util.Var.apply[T](init:T,e:com\.twitter\.util\.Event[T]):com\.twitter\.util\.Var[T]* Var.apply(T, Event[T])]] */ object Witness { /** - * Create a [[Witness]] from an [[AtomicReference atomic reference]]. + * Create a [[Witness]] from an atomic reference. */ def apply[T](ref: AtomicReference[T]): Witness[T] = new Witness[T] { def notify(t: T): Unit = ref.set(t) @@ -448,7 +498,7 @@ object Witness { } /** - * Create a [[Witness]] from a [[Function1]]. + * Create a [[Witness]] from a `Function1`. */ def apply[T](f: T => Unit): Witness[T] = new Witness[T] { def notify(t: T): Unit = f(t) @@ -463,7 +513,7 @@ object Witness { /** * Create a [[Witness]] which keeps only a `java.lang.ref.WeakReference` - * to the [[Function1]]. + * to the `Function1`. * * @see [[Var.patch]] for example. */ @@ -481,7 +531,7 @@ object Witness { * Create a [[Witness]] which keeps only a `java.lang.ref.WeakReference` * to `u`. * - * @see [[Var.apply]] for example. + * @see $VarApplyScaladocLink for example. */ def weakReference[T](u: Updatable[T]): Witness[T] = new Witness[T] { private[this] val weakRef = new WeakReference(u) diff --git a/util-core/src/main/scala/com/twitter/util/ExecutorWorkQueueFiber.scala b/util-core/src/main/scala/com/twitter/util/ExecutorWorkQueueFiber.scala new file mode 100644 index 0000000000..a492fc9c35 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/ExecutorWorkQueueFiber.scala @@ -0,0 +1,34 @@ +package com.twitter.util + +import java.util.concurrent.Executor + +/** + * A [[WorkQueueFiber]] implementation that submits its work queue to the underlying + * executor. + * + * This is intended to be used with a thread pool executor that manages task execution + * across threads. The ExecutorWorkQueueFiber pairs up a thread pool and a WorkQueueFiber + * in order to act as a scheduler and handle all task scheduling responsibilities + * for a server. + * + * @param pool a task Executor that must either execute the work immediately or pass the + * work to another thread in order to avoid potential deadlocks. + * @param fiberMetrics telemetry used for tracking the execution of this fiber, and likely + * other fibers as well. + */ +private[twitter] final class ExecutorWorkQueueFiber( + pool: Executor, + metrics: WorkQueueFiber.FiberMetrics) + extends WorkQueueFiber(metrics) { + + /** + * Submit the work to the underlying executor. + */ + override protected def schedulerSubmit(r: Runnable): Unit = pool.execute(r) + + /** + * No need to flush. Thread pools usually do not have per thread queues, and any that do + * should have some sort of work stealing mechanism to pass tasks around between threads. + */ + override protected def schedulerFlush(): Unit = () +} diff --git a/util-core/src/main/scala/com/twitter/util/Fiber.scala b/util-core/src/main/scala/com/twitter/util/Fiber.scala new file mode 100644 index 0000000000..ed2db6275d --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/Fiber.scala @@ -0,0 +1,55 @@ +package com.twitter.util + +import com.twitter.concurrent.Scheduler + +/** + * A Fiber receives runnable tasks and submits them to a Scheduler for execution. + */ +private[twitter] abstract class Fiber { + + /** Submit work to the Fiber */ + def submitTask(r: FiberTask): Unit + + /** Flush outstanding tasks */ + def flush(): Unit +} + +private[twitter] object Fiber extends Fiber { + + // Global default fiber which solely submits tasks to the global Scheduler + val Global: Fiber = new Fiber { + override def submitTask(r: FiberTask): Unit = { + Scheduler.submit(r) + } + override def flush(): Unit = { + Scheduler.flush() + } + } + + // Create fiber that captures the current Scheduler to avoid the volatile lookup on each + // submission + def newCachedSchedulerFiber(): Fiber = new Fiber { + // Just cache it so we don't have to keep looking at the volatile var each time. + private val scheduler = Scheduler() + def submitTask(r: FiberTask): Unit = { + scheduler.submit(r) + } + override def flush(): Unit = { + scheduler.flush() + } + } + + def let[T](fiber: Fiber)(f: => T): T = { + val oldCtx = Local.save() + val newCtx = oldCtx.setFiber(fiber) + Local.restore(newCtx) + try f + finally Local.restore(oldCtx) + } + + // Submit task to the fiber stored in the Local Context + def submitTask(r: FiberTask): Unit = Local.save().fiber.submitTask(r) + + // Flush fiber stored in the Local Context + def flush(): Unit = Local.save().fiber.flush() +} diff --git a/util-core/src/main/scala/com/twitter/util/FiberTask.scala b/util-core/src/main/scala/com/twitter/util/FiberTask.scala new file mode 100644 index 0000000000..909a4537fc --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/FiberTask.scala @@ -0,0 +1,19 @@ +package com.twitter.util + +/** + * Work for a `Fiber` to do. + * + * This is indistinguishable from a `Runnable` other than it is an abstract class. + * We do this so that in the case of the JVM not being able to optimize away the + * dynamic invocation we get v-table invocations which are constant time instead + * of i-table invocations which require a linear search through the dynamic types + * interfaces. + */ +private[twitter] abstract class FiberTask extends Runnable { + // From the `Runnable` interface + final def run(): Unit = doRun() + + // make this an abstract class method so we get v-table performance + // when invoked directly. + def doRun(): Unit +} diff --git a/util-core/src/main/scala/com/twitter/util/Function.scala b/util-core/src/main/scala/com/twitter/util/Function.scala index 7100f09aef..6b7b449684 100644 --- a/util-core/src/main/scala/com/twitter/util/Function.scala +++ b/util-core/src/main/scala/com/twitter/util/Function.scala @@ -51,7 +51,7 @@ object Function { def ofCallable[A](c: Callable[A]): () => A = () => c.call() /** - * Creates a [[Function]] of `T` to `R` from a [[JavaFunction]]. + * Creates a [[Function]] of `T` to `R` from a [[com.twitter.function.JavaFunction]]. * * Allows for better interop with Scala from Java 8 using lambdas. * @@ -99,7 +99,7 @@ object Function { } /** - * Creates a [[Function]] of `T` to `Unit` from a [[JavaConsumer]]. + * Creates a [[Function]] of `T` to `Unit` from a [[com.twitter.function.JavaConsumer]]. * * Allows for better interop with Scala from Java 8 using lambdas. * @@ -123,7 +123,8 @@ object Function { } /** - * Creates an [[ExceptionalFunction]] of `T` to `R` from an [[ExceptionalJavaFunction]]. + * Creates an [[ExceptionalFunction]] of `T` to `R` from an + * [[com.twitter.function.ExceptionalJavaFunction]]. * * Allows for better interop with Scala from Java 8 using lambdas. * @@ -143,7 +144,8 @@ object Function { } /** - * Creates an [[ExceptionalFunction]] of `T` to `Unit` from an [[ExceptionalJavaConsumer]]. + * Creates an [[ExceptionalFunction]] of `T` to `Unit` from an + * [[com.twitter.function.ExceptionalJavaConsumer]]. * * Allows for better interop with Scala from Java 8 using lambdas. * @@ -163,7 +165,8 @@ object Function { } /** - * Creates an [[ExceptionalFunction0]] to `T` from an [[ExceptionalSupplier]]. + * Creates an [[ExceptionalFunction0]] to `T` from an + * [[com.twitter.function.ExceptionalSupplier]]. * * Allows for better interop with Scala from Java 8 using lambdas. * diff --git a/util-core/src/main/scala/com/twitter/util/Future.scala b/util-core/src/main/scala/com/twitter/util/Future.scala index 9e8f88e9a5..fe334187c0 100644 --- a/util-core/src/main/scala/com/twitter/util/Future.scala +++ b/util-core/src/main/scala/com/twitter/util/Future.scala @@ -1,10 +1,16 @@ package com.twitter.util -import com.twitter.concurrent.{Offer, Tx} -import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReference} -import java.util.concurrent.{CancellationException, TimeUnit, Future => JavaFuture} +import com.twitter.concurrent.Offer +import com.twitter.concurrent.Tx +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.atomic.AtomicReference +import java.util.concurrent.CancellationException +import java.util.concurrent.CompletableFuture +import java.util.concurrent.{Future => JavaFuture} +import java.util.function.BiConsumer import java.util.{List => JList} -import scala.collection.mutable +import scala.collection.compat.immutable.ArraySeq +import scala.collection.{Seq => AnySeq} import scala.runtime.NonLocalReturnControl import scala.util.control.NoStackTrace @@ -12,6 +18,21 @@ import scala.util.control.NoStackTrace * @see [[Futures]] for Java-friendly APIs. * @see The [[https://twitter.github.io/finagle/guide/Futures.html user guide]] * on concurrent programming with Futures. + * + * @define FutureJoinScaladocLink + * [[[com.twitter.util.Future.join[A](fs:Seq[com\.twitter\.util\.Future[A]]):com\.twitter\.util\.Future[Unit]* join(Seq[A])]]] + * + * @define FuturesJoinScaladocLink + * [[com.twitter.util.Futures.join[A,B](a:com\.twitter\.util\.Future[A],b:com\.twitter\.util\.Future[B]):com\.twitter\.util\.Future[(A,B)]* Futures.join(Future[A],Future[B])]] + * + * @define FutureCollectScaladocLink + * [[[com.twitter.util.Future.collect[A](fs:Seq[com\.twitter\.util\.Future[A]]):com\.twitter\.util\.Future[Seq[A]]* collect(Seq[Future[A]])]]] + * + * @define FutureCollectToTryScaladocLink + * [[[[com.twitter.util.Future.collectToTry[A](fs:Seq[com\.twitter\.util\.Future[A]]):com\.twitter\.util\.Future[Seq[com\.twitter\.util\.Try[A]]]* collectToTry(Seq[Future[A]])]]]] + * + * @define FutureSelectScaladocLink + * [[[com.twitter.util.Future.select[A](fs:Seq[com\.twitter\.util\.Future[A]]):com\.twitter\.util\.Future[(com\.twitter\.util\.Try[A],Seq[com\.twitter\.util\.Future[A]])]* select(Seq.MethodPerEndpoint)]]] */ object Future { val DEFAULT_TIMEOUT: Duration = Duration.Top @@ -40,48 +61,38 @@ object Future { /** A successfully satisfied constant `Future` of `false` */ val False: Future[Boolean] = new ConstFuture(Return.False) - /** - * A failed `Future` analogous to `Predef.???`. - */ - val ??? : Future[Nothing] = - Future.exception(new NotImplementedError("an implementation is missing")) - - private val SomeReturnUnit = Some(Return.Unit) private val NotApplied: Future[Nothing] = new NoFuture private val AlwaysNotApplied: Any => Future[Nothing] = scala.Function.const(NotApplied) - private val toUnit: Try[Any] => Future[Unit] = { - case Return(_) => Future.Unit - case t => Future.const(t.asInstanceOf[Try[Unit]]) + private val tryToUnit: Try[Any] => Try[Unit] = { + case t: Return[_] => Try.Unit + case t => t.asInstanceOf[Try[Unit]] } - private val toVoid: Try[Any] => Future[Void] = { - case Return(_) => Future.Void - case t => Future.const(t.asInstanceOf[Try[Void]]) + private val tryToVoid: Try[Any] => Try[Void] = { + case t: Return[_] => Try.Void + case t => t.asInstanceOf[Try[Void]] } private val AlwaysMasked: PartialFunction[Throwable, Boolean] = { case _ => true } private val toTuple2Instance: (Any, Any) => (Any, Any) = Tuple2.apply private def toTuple2[A, B]: (A, B) => (A, B) = toTuple2Instance.asInstanceOf[(A, B) => (A, B)] - private val toValueInstance: Any => Future[Any] = Future.value - private def toValue[T]: T => Future[T] = toValueInstance.asInstanceOf[T => Future[T]] - - private val lowerFromTryInstance: Any => Future[Any] = { - case t: Try[_] => Future.const(t) - } + private val liftToTryInstance: Any => Try[Any] = Return(_) + private def liftToTry[T]: T => Try[T] = liftToTryInstance.asInstanceOf[T => Try[T]] - private def lowerFromTry[A, T]: A => Future[T] = - lowerFromTryInstance.asInstanceOf[A => Future[T]] + private val flattenTryInstance: Try[Try[Any]] => Try[Any] = _.flatten + private def flattenTry[A, T](implicit ev: A => Try[T]): Try[A] => Try[T] = + flattenTryInstance.asInstanceOf[Try[A] => Try[T]] - private val toTxInstance: Try[Any] => Future[Tx[Try[Any]]] = + private val toTxInstance: Try[Any] => Try[Tx[Try[Any]]] = res => { val tx = new Tx[Try[Any]] { def ack(): Future[Tx.Result[Try[Any]]] = Future.value(Tx.Commit(res)) def nack(): Unit = () } - Future.value(tx) + Return(tx) } - private def toTx[A]: Try[A] => Future[Tx[Try[A]]] = - toTxInstance.asInstanceOf[Try[A] => Future[Tx[Try[A]]]] + private def toTx[A]: Try[A] => Try[Tx[Try[A]]] = + toTxInstance.asInstanceOf[Try[A] => Try[Tx[Try[A]]]] private val emptySeqInstance: Future[Seq[Any]] = Future.value(Seq.empty) private def emptySeq[A]: Future[Seq[A]] = emptySeqInstance.asInstanceOf[Future[Seq[A]]] @@ -93,6 +104,12 @@ object Future { private[this] val RaiseException = new Exception with NoStackTrace @inline private final def raiseException = RaiseException + /** + * A failed `Future` analogous to `Predef.???`. + */ + def ??? : Future[Nothing] = + Future.exception(new NotImplementedError("an implementation is missing")) + /** * Creates a satisfied `Future` from a [[Try]]. * @@ -159,8 +176,6 @@ object Future { case nlrc: NonLocalReturnControl[_] => Future.exception(new FutureNonLocalReturnControl(nlrc)) } - def unapply[A](f: Future[A]): Option[Try[A]] = f.poll - // The thread-safety for `outer` is guaranteed by synchronizing on `this`. private[this] final class MonitoredPromise[A](private[this] var outer: Promise[A]) extends Monitor @@ -230,7 +245,7 @@ object Future { // but because these are both Function1 of different types, this cannot be done. Because // the respond handler is invoked for each Future in the sequence, and the interrupt handler is // only set once, we prefer to mix in Function1. - private[this] class JoinPromise[A](fs: Seq[Future[A]], size: Int) + private[this] class JoinPromise[A](fs: Iterable[Future[A]], size: Int) extends Promise[Unit] with (Try[A] => Unit) { @@ -261,12 +276,12 @@ object Future { * * @param fs a sequence of Futures * - * @see [[Futures.join]] for a Java friendly API. - * @see [[collect]] if you want to be able to see the results of each `Future`. - * @see [[collectToTry]] if you want to be able to see the results of each - * `Future` regardless of if they succeed or fail. + * @see $FuturesJoinScaladocLink for a Java friendly API. + * @see $FutureCollectScaladocLink if you want to be able to see the results of each `Future`. + * @see $FutureCollectToTryScaladocLink if you want to be able to see the results of each + * `Future` regardless of if they succeed or fail. */ - def join[A](fs: Seq[Future[A]]): Future[Unit] = { + def join[A](fs: AnySeq[Future[A]]): Future[Unit] = { if (fs.isEmpty) Future.Unit else { val size = fs.size @@ -300,7 +315,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( meths foreach println ' - */ + """*/ /** * Join 2 futures. The returned future is complete when all @@ -317,9 +332,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * do. */ def join[A, B, C](a: Future[A], b: Future[B], c: Future[C]): Future[(A, B, C)] = - join(Seq(a, b, c)).map { _ => - (Await.result(a), Await.result(b), Await.result(c)) - } + join(Seq(a, b, c)).map { _ => (Await.result(a), Await.result(b), Await.result(c)) } /** * Join 4 futures. The returned future is complete when all @@ -1034,10 +1047,10 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * * @param fs a java.util.List of Futures * - * @see [[Futures.join]] for a Java friendly API. - * @see [[collect]] if you want to be able to see the results of each `Future`. - * @see [[collectToTry]] if you want to be able to see the results of each - * `Future` regardless of if they succeed or fail. + * @see $FuturesJoinScaladocLink for a Java friendly API. + * @see $FutureCollectScaladocLink if you want to be able to see the results of each `Future`. + * @see $FutureCollectToTryScaladocLink if you want to be able to see the results of each + * `Future` regardless of if they succeed or fail. */ def join[A](fs: JList[Future[A]]): Future[Unit] = Futures.join(fs) @@ -1053,7 +1066,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * {{{ * // will return a Future of `Seq(2, 3, 4)` * Future.traverseSequentially(Seq(1, 2, 3)) { i => - * Future(i + 1) + * Future.value(i + 1) * } * }}} * @@ -1068,12 +1081,16 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( } yield results :+ nextResult } - private[this] final class CollectPromise[A](fs: Seq[Future[A]]) + private[this] final class CollectPromise[A](fs: Iterable[Future[A]]) extends Promise[Seq[A]] with Promise.InterruptHandler { - private[this] val results = new mutable.ArraySeq[A](fs.size) - private[this] val count = new AtomicInteger(results.size) + private[this] val results = Array.ofDim[Any](fs.size) + private[this] def setResultsAsValue() = { + val seq = ArraySeq.unsafeWrapArray(results).asInstanceOf[ArraySeq[A]] + setValue(seq) + } + private[this] val count = new AtomicInteger(fs.size) // Respond handler. It's safe to write into different array cells concurrently. // We guarantee the thread writing a last value will observe all previous writes @@ -1082,7 +1099,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( def collectTo(index: Int): Try[A] => Unit = { case Return(a) => results(index) = a - if (count.decrementAndGet() == 0) setValue(results) + if (count.decrementAndGet() == 0) setResultsAsValue() case t @ Throw(_) => updateIfEmpty(t.cast[Seq[A]]) } @@ -1103,12 +1120,12 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * @param fs a sequence of Futures * @return a `Future[Seq[A]]` containing the collected values from fs. * - * @see [[collectToTry]] if you want to be able to see the results of each + * @see $FutureCollectToTryScaladocLink if you want to be able to see the results of each * `Future` regardless of if they succeed or fail. - * @see [[join]] if you are not interested in the results of the individual + * @see $FutureJoinScaladocLink if you are not interested in the results of the individual * `Futures`, only when they are complete. */ - def collect[A](fs: Seq[Future[A]]): Future[Seq[A]] = + def collect[A](fs: AnySeq[Future[A]]): Future[Seq[A]] = if (fs.isEmpty) emptySeq else { val result = new CollectPromise[A](fs) @@ -1135,9 +1152,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( if (fs.isEmpty) emptyMap else { val (keys, values) = fs.toSeq.unzip - Future.collect(values).map { seq => - keys.zip(seq)(scala.collection.breakOut): Map[A, B] - } + Future.collect(values).map { seq => keys.iterator.zip(seq.iterator).toMap: Map[A, B] } } /** @@ -1161,13 +1176,24 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * @param fs a sequence of Futures * @return a `Future[Seq[Try[A]]]` containing the collected values from fs. */ - def collectToTry[A](fs: Seq[Future[A]]): Future[Seq[Try[A]]] = - Future.collect { - val seq = new mutable.ArrayBuffer[Future[Try[A]]](fs.size) + def collectToTry[A](fs: AnySeq[Future[A]]): Future[Seq[Try[A]]] = { + //unroll cases 0 and 1 + if (fs.isEmpty) Nil + else { val iterator = fs.iterator - while (iterator.hasNext) seq += iterator.next().liftToTry - seq + val h = iterator.next() + if (iterator.hasNext) { + val buf = Vector.newBuilder[Future[Try[A]]] + buf.sizeHint(fs.size) + buf += h.liftToTry + buf += iterator.next().liftToTry + while (iterator.hasNext) buf += iterator.next().liftToTry + Future.collect(buf.result()) + } else { + Future.collect(List(h.liftToTry)) + } } + } /** * Collect the results from the given futures into a new future of java.util.List[Try[A]]. @@ -1214,15 +1240,16 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( a.respond { t => if (!p.isDefined) { val filtered = { - val seq = new mutable.ArrayBuffer[Future[A]](size - 1) + val buf = Vector.newBuilder[Future[A]] + buf.sizeHint(size - 1) var j = 0 while (j < size) { val (_, fi) = as(j) if (fi ne f) - seq += fi + buf += fi j += 1 } - seq + buf.result() } p.updateIfEmpty(Return(t -> filtered)) @@ -1243,7 +1270,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * * @param fs cannot be empty * - * @see [[select]] which can be an easier API to use. + * @see $FutureSelectScaladocLink which can be an easier API to use. */ def selectIndex[A](fs: IndexedSeq[Future[A]]): Future[Int] = if (fs.isEmpty) { @@ -1312,9 +1339,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( def whileDo[A](p: => Boolean)(f: => Future[A]): Future[Unit] = { def loop(): Future[Unit] = { if (p) { - f.flatMap { _ => - loop() - } + f.flatMap { _ => loop() } } else Future.Unit } @@ -1324,13 +1349,12 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( case class NextThrewException(cause: Throwable) extends IllegalArgumentException("'next' threw an exception", cause) - private class Each[A](next: => Future[A], body: A => Unit) - extends (Try[A] => Future[Nothing]) { + private class Each[A](next: => Future[A], body: A => Unit) extends (Try[A] => Future[Nothing]) { def apply(t: Try[A]): Future[Nothing] = t match { case Return(a) => body(a) go() - case t@Throw(_) => + case t @ Throw(_) => Future.const(t.cast[Nothing]) } def go(): Future[Nothing] = { @@ -1352,17 +1376,18 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( def parallel[A](n: Int)(f: => Future[A]): Seq[Future[A]] = { var i = 0 - val result = new mutable.ArrayBuffer[Future[A]](n) + val buf = Vector.newBuilder[Future[A]] + buf.sizeHint(n) while (i < n) { - result += f + buf += f i += 1 } - result + buf.result() } /** * Creates a "batched" Future that, given a function - * `Seq[In] => Future[Seq[Out]]`, returns a `In => Future[Out]` interface + * `scala.collection.Seq[In] => Future[Seq[Out]]`, returns a `In => Future[Out]` interface * that batches the underlying asynchronous operations. Thus, one can * incrementally submit tasks to be performed when the criteria for batch * flushing is met. @@ -1399,14 +1424,44 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( timeThreshold: Duration = Duration.Top, sizePercentile: => Float = 1.0f )( - f: Seq[In] => Future[Seq[Out]] + f: AnySeq[In] => Future[Seq[Out]] )( implicit timer: Timer ): Batcher[In, Out] = { new Batcher[In, Out]( - new BatchExecutor[In, Out](sizeThreshold, timeThreshold, sizePercentile, f) + new BatchExecutor[In, Out]( + sizeThreshold, + timeThreshold, + sizePercentile, + f.compose(it => it.toSeq)) ) } + + /** + * Convert a Java native [[CompletableFuture]] to a Twitter [[Future]]. + * + * @note The Twitter Future's cancellation will be propagated to the underlying [[CompletableFuture]] + * @see [[https://download.oracle.com/javase/6/docs/api/java/util/concurrent/Future.html#cancel(boolean)]] + */ + def fromCompletableFuture[T](completableFuture: CompletableFuture[T]): Future[T] = { + val p = new Promise[T] + + // ensure that we propagate cancellation/errors from the Promise to the CompletableFuture. + // without this propagation, work contained within the CompletableFuture chain will continue, + // despite the Promise completing in error. we expect the CompletableFuture chain to reflect + // that an error has occurred via cancellation and stop work that would otherwise be wasted. + p.setInterruptHandler { + case _ => completableFuture.cancel(true) + } + + // ensure the state of the CompletableFuture propagates back to the Promise on completion + completableFuture.whenComplete { + case (_, thrown: Throwable) if thrown != null => p.setException(thrown) + case (t, _) => p.setValue(t) + } + + p + } } /** @@ -1427,6 +1482,23 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * @see The [[https://twitter.github.io/finagle/guide/Futures.html user guide]] * on concurrent programming with Futures. * @see [[Futures]] for Java-friendly APIs. + * + * @define callbacks + * {{{ + * import com.twitter.util.Future + * def callbacks(result: Future[Int]): Future[Int] = + * result.onSuccess { i => + * println(i) + * }.onFailure { e => + * println(e.getMessage) + * }.ensure { + * println("always printed") + * } + * }}} + * + * @define awaitresult + * // Await.result blocks the current thread, + * // don't use it except for tests. */ abstract class Future[+A] extends Awaitable[A] { self => @@ -1441,8 +1513,23 @@ abstract class Future[+A] extends Awaitable[A] { self => * libraries). Otherwise, it is a best practice to use one of the * alternatives ([[onSuccess]], [[onFailure]], etc.). * - * @note this should be used for side-effects. + * @example + * {{{ + * import com.twitter.util.{Await, Future, Return, Throw} + * val f1: Future[Int] = Future.value(1) + * val f2: Future[Int] = Future.exception(new Exception("boom!")) + * val printing: Future[Int] => Future[Int] = x => x.respond { + * case Return(_) => println("Here's a Return") + * case Throw(_) => println("Here's an Exception") + * } + * Await.result(printing(f1)) + * // Prints side-effect "Here's a Return" and then returns Future value "1" + * Await.result(printing(f2)) + * // Prints side-effect "Here's an Exception" and then throws java.lang.Exception: boom! + * $awaitresult + * }}} * + * @note this should be used for side-effects. * @param k the side-effect to apply when the computation completes. * The value of the input to `k` will be the result of the * computation to this future. @@ -1463,29 +1550,25 @@ abstract class Future[+A] extends Awaitable[A] { self => * The returned `Future` will be satisfied when this, * the original future, is done. * - * @note this should be used for side-effects. + * @example + * $callbacks + * {{{ + * val a = Future.value(1) + * callbacks(a) // prints "1" and then "always printed" + * }}} * + * @note this should be used for side-effects. * @param f the side-effect to apply when the computation completes. * @see [[respond]] if you need the result of the computation for * usage in the side-effect. */ - def ensure(f: => Unit): Future[A] = respond { _ => - f - } + def ensure(f: => Unit): Future[A] = respond { _ => f } /** * Is the result of the Future available yet? */ def isDefined: Boolean = poll.isDefined - /** - * Checks whether a Unit-typed Future is done. By - * convention, futures of type Future[Unit] are used - * for signalling. - */ - def isDone(implicit ev: this.type <:< Future[Unit]): Boolean = - ev(this).poll == Future.SomeReturnUnit - /** * Polls for an available result. If the Future has been * satisfied, returns Some(result), otherwise None. @@ -1561,7 +1644,7 @@ abstract class Future[+A] extends Awaitable[A] { self => * * ''Note'': On timeout, the underlying future is interrupted. */ - def raiseWithin(timeout: Duration, exc: Throwable)(implicit timer: Timer): Future[A] = + def raiseWithin(timeout: Duration, exc: => Throwable)(implicit timer: Timer): Future[A] = raiseWithin(timer, timeout, exc) /** @@ -1569,7 +1652,7 @@ abstract class Future[+A] extends Awaitable[A] { self => * * ''Note'': On timeout, the underlying future is interrupted. */ - def raiseWithin(timer: Timer, timeout: Duration, exc: Throwable): Future[A] = { + def raiseWithin(timer: Timer, timeout: Duration, exc: => Throwable): Future[A] = { if (timeout == Duration.Top || isDefined) return this @@ -1681,9 +1764,23 @@ abstract class Future[+A] extends Awaitable[A] { self => * When this future completes, run `f` on that completed result * whether or not this computation was successful. * - * The returned `Future` will be satisfied when this, + * The returned `Future` will be satisfied when `this`, * the original future, and `f` are done. * + * @example + * {{{ + * import com.twitter.util.{Await, Future, Return, Throw} + * val f1: Future[Int] = Future.value(1) + * val f2: Future[Int] = Future.exception(new Exception("boom!")) + * val transforming: Future[Int] => Future[String] = x => x.transform { + * case Return(i) => Future.value(i.toString) + * case Throw(e) => Future.value(e.getMessage) + * } + * Await.result(transforming(f1)) // String = 1 + * Await.result(transforming(f2)) // String = "boom!" + * $awaitresult + * }}} + * * @see [[respond]] for purely side-effecting callbacks. * @see [[map]] and [[flatMap]] for dealing strictly with successful * computations. @@ -1693,13 +1790,57 @@ abstract class Future[+A] extends Awaitable[A] { self => */ def transform[B](f: Try[A] => Future[B]): Future[B] + /** + * When this future completes, run `f` on that completed result + * whether or not this computation was successful. + * + * @note This method is similar to `transform`, but the transformation is + * applied without introducing an intermediate future, which leads + * to fewer allocations. The promise is satisfied directly instead + * of wrapping a Future around a Try. + * + * The returned `Future` will be satisfied when `this`, + * the original future, is done. + * + * @see [[respond]] for purely side-effecting callbacks. + * @see [[map]] and [[flatMap]] for dealing strictly with successful + * computations. + * @see [[handle]] and [[rescue]] for dealing strictly with exceptional + * computations. + * @see [[transformedBy]] for a Java friendly API. + * @see [[transform]] for transformations that return `Future`. + */ + protected def transformTry[B](f: Try[A] => Try[B]): Future[B] + /** * If this, the original future, succeeds, run `f` on the result. * * The returned result is a Future that is satisfied when the original future * and the callback, `f`, are done. + * + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f: Future[Int] = Future.value(1) + * val newf: Future[Int] = f.flatMap { x => + * Future.value(x + 10) + * } + * Await.result(newf) // 11 + * $awaitresult + * }}} + * * If the original future fails, this one will also fail, without executing `f` - * and preserving the failed computation of `this`. + * and preserving the failed computation of `this` + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f: Future[Int] = Future.exception(new Exception("boom!")) + * val newf: Future[Int] = f.flatMap { x => + * println("I'm being executed") // won't print + * Future.value(x + 10) + * } + * Await.result(newf) // throws java.lang.Exception: boom! + * }}} * * @see [[map]] */ @@ -1729,11 +1870,22 @@ abstract class Future[+A] extends Awaitable[A] { self => * * This is the equivalent of [[flatMap]] for failed computations. * + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f1: Future[Int] = Future.exception(new Exception("boom1!")) + * val f2: Future[Int] = Future.exception(new Exception("boom2!")) + * val newf: Future[Int] => Future[Int] = x => x.rescue { + * case e: Exception if e.getMessage == "boom1!" => Future.value(1) + * } + * Await.result(newf(f1)) // 1 + * Await.result(newf(f2)) // throws java.lang.Exception: boom2! + * $awaitresult + * }}} + * * @see [[handle]] */ - def rescue[B >: A]( - rescueException: PartialFunction[Throwable, Future[B]] - ): Future[B] = + def rescue[B >: A](rescueException: PartialFunction[Throwable, Future[B]]): Future[B] = transform { case Throw(t) => val result = rescueException.applyOrElse(t, Future.AlwaysNotApplied) @@ -1754,21 +1906,40 @@ abstract class Future[+A] extends Awaitable[A] { self => * * The returned result is a Future that is satisfied when the original future * and the callback, `f`, are done. + * + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f: Future[Int] = Future.value(1) + * val newf: Future[Int] = f.map { x => + * x + 10 + * } + * Await.result(newf) // 11 + * $awaitresult + * }}} + * * If the original future fails, this one will also fail, without executing `f` - * and preserving the failed computation of `this`. + * and preserving the failed computation of `this` + * + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f: Future[Int] = Future.exception(new Exception("boom!")) + * val newf: Future[Int] = f.map { x => + * println("I'm being executed") // won't print + * x + 10 + * } + * Await.result(newf) // throws java.lang.Exception: boom! + * }}} * * @see [[flatMap]] for computations that return `Future`s. * @see [[onSuccess]] for side-effecting chained computations. */ def map[B](f: A => B): Future[B] = - transform { - case Return(r) => Future { f(r) } - case t: Throw[_] => Future.const[B](t.cast[B]) - } + transformTry(_.map(f)) - def filter(p: A => Boolean): Future[A] = transform { x: Try[A] => - Future.const(x.filter(p)) - } + def filter(p: A => Boolean): Future[A] = + transformTry(_.filter(p)) def withFilter(p: A => Boolean): Future[A] = filter(p) @@ -1778,6 +1949,13 @@ abstract class Future[+A] extends Awaitable[A] { self => * * @note this should be used for side-effects. * + * @example + * $callbacks + * {{{ + * val a = Future.value(1) + * callbacks(a) // prints "1" and then "always printed" + * }}} + * * @return chained Future * @see [[flatMap]] and [[map]] to produce a new `Future` from the result of * the computation. @@ -1800,6 +1978,13 @@ abstract class Future[+A] extends Awaitable[A] { self => * `future.onFailure { case NonFatal(e) => ... }` when the Throwable * is "fatal". * + * @example + * $callbacks + * {{{ + * val b = Future.exception(new Exception("boom!")) + * callbacks(b) // prints "boom!" and then "always printed" + * }}} + * * @return chained Future * @see [[handle]] and [[rescue]] to produce a new `Future` from the result of * the computation. @@ -1855,6 +2040,19 @@ abstract class Future[+A] extends Awaitable[A] { self => * * This is the equivalent of [[map]] for failed computations. * + * @example + * {{{ + * import com.twitter.util.{Await, Future} + * val f1: Future[Int] = Future.exception(new Exception("boom1!")) + * val f2: Future[Int] = Future.exception(new Exception("boom2!")) + * val newf: Future[Int] => Future[Int] = x => x.handle { + * case e: Exception if e.getMessage == "boom1!" => 1 + * } + * Await.result(newf(f1)) // 1 + * Await.result(newf(f2)) // throws java.lang.Exception: boom2! + * $awaitresult + * }}} + * * @see [[rescue]] */ def handle[B >: A](rescueException: PartialFunction[Throwable, B]): Future[B] = rescue { @@ -1872,12 +2070,8 @@ abstract class Future[+A] extends Awaitable[A] { self => val p = Promise.interrupts[U](other, this) val a = Promise.attached(other) val b = Promise.attached(this) - a.respond { t => - if (p.updateIfEmpty(t)) b.detach() - } - b.respond { t => - if (p.updateIfEmpty(t)) a.detach() - } + a.respond { t => if (p.updateIfEmpty(t)) b.detach() } + b.respond { t => if (p.updateIfEmpty(t)) a.detach() } p } @@ -1923,7 +2117,7 @@ abstract class Future[+A] extends Awaitable[A] { self => tx match { case Return(_) if !isFirst && get.isReturn => p.setValue(fn(Await.result(Future.this), Await.result(other))) - case t@Throw(_) if isFirst || get.isReturn => + case t @ Throw(_) if isFirst || get.isReturn => p.update(t.cast[C]) case _ => } @@ -1940,14 +2134,14 @@ abstract class Future[+A] extends Awaitable[A] { self => * * @note failed futures will remain as is. */ - def unit: Future[Unit] = transform(Future.toUnit) + def unit: Future[Unit] = transformTry(Future.tryToUnit) /** * Convert this `Future[A]` to a `Future[Void]` by discarding the result. * * @note failed futures will remain as is. */ - def voided: Future[Void] = transform(Future.toVoid) + def voided: Future[Void] = transformTry(Future.tryToVoid) /** * Send updates from this Future to the other. @@ -1976,41 +2170,55 @@ abstract class Future[+A] extends Awaitable[A] { self => * The offer is activated when the future is satisfied. */ def toOffer: Offer[Try[A]] = new Offer[Try[A]] { - def prepare(): Future[Tx[Try[A]]] = transform(Future.toTx) + def prepare(): Future[Tx[Try[A]]] = transformTry(Future.toTx) } /** * Convert a Twitter Future to a Java native Future. This should * match the semantics of a Java Future as closely as possible to - * avoid issues with the way another API might use them. See: + * avoid issues with the way another API might use them. At the same time, + * its semantics should be similar to those of a Twitter Future, so it + * propagates cancellation. See: * - * http://download.oracle.com/javase/6/docs/api/java/util/concurrent/Future.html#cancel(boolean) + * https://download.oracle.com/javase/6/docs/api/java/util/concurrent/Future.html#cancel(boolean) */ - def toJavaFuture: JavaFuture[_ <: A] = { - new JavaFuture[A] { - private[this] val wasCancelled = new AtomicBoolean(false) - - def cancel(mayInterruptIfRunning: Boolean): Boolean = { - if (wasCancelled.compareAndSet(false, true)) - self.raise(new CancellationException) - true - } + def toJavaFuture: JavaFuture[_ <: A] = toCompletableFuture - def isCancelled: Boolean = wasCancelled.get - def isDone: Boolean = isCancelled || self.isDefined + /** + * Convert a Twitter Future to a Java native CompletableFuture. This should + * match the semantics of a Java Future as closely as possible to + * avoid issues with the way another API might use them. At the same time, + * its semantics should be similar to those of a Twitter Future, so it + * propagates cancellation. See: + * + * https://download.oracle.com/javase/6/docs/api/java/util/concurrent/Future.html#cancel(boolean) + */ + def toCompletableFuture[B >: A]: CompletableFuture[B] = { + val f = new CompletableFuture[B]() with (Try[A] => Unit) with BiConsumer[B, Throwable] { + // Complete the Java future from a Try + def apply(t: Try[A]): Unit = + t match { + case Return(result) => complete(result) + case Throw(t) => completeExceptionally(t) + } - def get(): A = { - if (isCancelled) - throw new CancellationException - Await.result(self) - } + // For propagating cancellation + def accept(r: B, t: Throwable): Unit = + if (t.isInstanceOf[CancellationException]) { + self.raise(t) + } + } - def get(time: Long, timeUnit: TimeUnit): A = { - if (isCancelled) - throw new CancellationException - Await.result(self, Duration.fromTimeUnit(time, timeUnit)) - } + poll match { + case Some(t) => f(t) // complete + case None => + this.respond(f) // propagate completion forward + f.whenComplete(f) // propagate cancellation backwards } + + // we have to return the original one because cancellation isn't propagated + // by CompletionStages + f } /** @@ -2029,7 +2237,7 @@ abstract class Future[+A] extends Awaitable[A] { self => * For example: * {{{ * val p = new Promise[Int]() - * p.setInterruptHandler { case x => println(s"interrupt handler for ${x.getClass}") } + * p.setInterruptHandler { case x => println(s"interrupt handler for \${x.getClass}") } * val f1: Future[Int] = p.mask { * case _: IllegalArgumentException => true * } @@ -2044,7 +2252,7 @@ abstract class Future[+A] extends Awaitable[A] { self => * * @see [[raise]] * @see [[masked]] - * @see [[interruptible()]] + * @see [[interruptible]] */ def mask(pred: PartialFunction[Throwable, Boolean]): Future[A] = { val p = Promise[A]() @@ -2075,7 +2283,7 @@ abstract class Future[+A] extends Awaitable[A] { self => * * @see [[raise]] * @see [[mask]] - * @see [[interruptible()]] + * @see [[interruptible]] */ def masked: Future[A] = mask(Future.AlwaysMasked) @@ -2086,22 +2294,44 @@ abstract class Future[+A] extends Awaitable[A] { self => */ def willEqual[B](that: Future[B]): Future[Boolean] = { this.transform { thisResult => - that.transform { thatResult => - Future.value(thisResult == thatResult) - } + that.transform { thatResult => Future.value(thisResult == thatResult) } } } /** * Returns the result of the computation as a `Future[Try[A]]`. + * + * @example + * {{{ + * import com.twitter.util.{Await, Future, Try} + * val fr: Future[Int] = Future.value(1) + * val ft: Future[Int] = Future.exception(new Exception("boom!")) + * val r: Future[Try[Int]] = fr.liftToTry + * val t: Future[Try[Int]] = ft.liftToTry + * Await.result(r) // Return(1) + * Await.result(t) // Throw(java.lang.Exception: boom!) + * $awaitresult + * }}} */ - def liftToTry: Future[Try[A]] = transform(Future.toValue) + def liftToTry: Future[Try[A]] = transformTry(Future.liftToTry) /** * Lowers a `Future[Try[T]]` into a `Future[T]`. + * + * @example + * {{{ + * import com.twitter.util.{Await, Future, Return, Throw, Try} + * val fr: Future[Try[Int]] = Future.value(Return(1)) + * val ft: Future[Try[Int]] = Future.value(Throw(new Exception("boom!"))) + * val r: Future[Int] = fr.lowerFromTry + * val t: Future[Int] = ft.lowerFromTry + * Await.result(r) // 1 + * Await.result(t) // throws java.lang.Exception: boom! + * $awaitresult + * }}} */ def lowerFromTry[B](implicit ev: A <:< Try[B]): Future[B] = - flatMap(Future.lowerFromTry) + transformTry(Future.flattenTry) /** * Makes a derivative `Future` which will be satisfied with the result diff --git a/util-core/src/main/scala/com/twitter/util/FutureEventListener.scala b/util-core/src/main/scala/com/twitter/util/FutureEventListener.scala index 4b8f2a2699..deb533eb4a 100644 --- a/util-core/src/main/scala/com/twitter/util/FutureEventListener.scala +++ b/util-core/src/main/scala/com/twitter/util/FutureEventListener.scala @@ -23,4 +23,3 @@ trait FutureEventListener[T] { */ def onFailure(cause: Throwable): Unit } - diff --git a/util-core/src/main/scala/com/twitter/util/FuturePool.scala b/util-core/src/main/scala/com/twitter/util/FuturePool.scala index acfa8c414f..ad66a238e2 100644 --- a/util-core/src/main/scala/com/twitter/util/FuturePool.scala +++ b/util-core/src/main/scala/com/twitter/util/FuturePool.scala @@ -1,8 +1,9 @@ package com.twitter.util +import com.twitter.concurrent.ForkingScheduler import com.twitter.concurrent.NamedPoolThreadFactory import java.util.concurrent.atomic.AtomicBoolean -import java.util.concurrent._ +import java.util.concurrent.{Future => _, _} import scala.runtime.NonLocalReturnControl /** @@ -39,6 +40,13 @@ trait FuturePool { * @return -1 if this [[FuturePool]] implementation doesn't support this statistic. */ def numCompletedTasks: Long = -1L + + /** + * The approximate number of tasks that are pending (typically waiting in the queue). + * + * @return -1 if this [[FuturePool]] implementation doesn't support this statistic. + */ + def numPendingTasks: Long = -1 } /** @@ -69,13 +77,13 @@ object FuturePool { override def toString: String = "FuturePool.immediatePool" } - private lazy val defaultExecutor = Executors.newCachedThreadPool( + private[twitter] lazy val defaultExecutor = Executors.newCachedThreadPool( new NamedPoolThreadFactory("UnboundedFuturePool", makeDaemons = true) ) /** * The default future pool, using a cached threadpool, provided by - * [[java.util.concurrent.Executors.newCachedThreadPool]]. Note + * `java.util.concurrent.Executors.newCachedThreadPool`. Note * that this is intended for IO concurrency; computational * parallelism typically requires special treatment. If an interrupt * is raised on a returned Future and the work has started, the worker @@ -89,7 +97,7 @@ object FuturePool { /** * The default future pool, using a cached threadpool, provided by - * [[java.util.concurrent.Executors.newCachedThreadPool]]. Note + * `java.util.concurrent.Executors.newCachedThreadPool`. Note * that this is intended for IO concurrency; computational * parallelism typically requires special treatment. If an interrupt * is raised on a returned Future and the work has started, an attempt @@ -118,64 +126,89 @@ class InterruptibleExecutorServiceFuturePool(executor: ExecutorService) */ class ExecutorServiceFuturePool protected[this] ( val executor: ExecutorService, - val interruptible: Boolean -) extends FuturePool { + val interruptible: Boolean) + extends FuturePool { def this(executor: ExecutorService) = this(executor, false) - def apply[T](f: => T): Future[T] = { - val runOk = new AtomicBoolean(true) - val p = new Promise[T] - val task = new Runnable { - private[this] val saved = Local.save() - - def run(): Unit = { - // Make an effort to skip work in the case the promise - // has been cancelled or already defined. - if (!runOk.compareAndSet(true, false)) - return + /** + * Not null if the current scheduler supports forking and indicates that + * `FuturePool`s should be redirected to it. If the provided executor is a + * `ThreadPoolExecutor`, the max concurrency of the forking scheduler is + * set to the size of the thread pool, maintaining a similar concurrency behavior + */ + private[this] lazy val forkingScheduler: ForkingScheduler = { + ForkingScheduler() + .filter(_.redirectFuturePools()) + .map { s => + executor match { + case e: ThreadPoolExecutor => + s.withMaxSyncConcurrency(e.getMaximumPoolSize, e.getQueue.remainingCapacity()) + case _ => + s + } + }.getOrElse(null) + } - val current = Local.save() - Local.restore(saved) + def apply[T](f: => T): Future[T] = { + val saved = Local.save() + val tracker = saved.resourceTracker + val applyF = + if (tracker eq None) () => f + else ResourceTracker.wrapAndMeasureUsage[T](f, tracker.get) + + if (forkingScheduler != null) forkingScheduler.fork(Future.value(applyF())) + else { + val runOk = new AtomicBoolean(true) + val p = new Promise[T] + val task = new Runnable { + def run(): Unit = { + // Make an effort to skip work in the case the promise + // has been cancelled or already defined. + if (!runOk.compareAndSet(true, false)) + return + + val current = Local.save() + if (current ne saved) Local.restore(saved) + + try p.updateIfEmpty(Try(applyF())) + catch { + case nlrc: NonLocalReturnControl[_] => + val fnlrc = new FutureNonLocalReturnControl(nlrc) + p.updateIfEmpty(Throw(fnlrc)) + throw fnlrc + case e: Throwable => + p.updateIfEmpty(Throw(new ExecutionException(e))) + throw e + } finally Local.restore(current) + } + } - try p.updateIfEmpty(Try(f)) + // This is safe: the only thing that can call task.run() is + // executor, the only thing that can raise an interrupt is the + // receiver of this value, which will then be fully initialized. + val javaFuture = + try executor.submit(task) catch { - case nlrc: NonLocalReturnControl[_] => - val fnlrc = new FutureNonLocalReturnControl(nlrc) - p.updateIfEmpty(Throw(fnlrc)) - throw fnlrc - case e: Throwable => - p.updateIfEmpty(Throw(new ExecutionException(e))) - throw e - } finally Local.restore(current) - } - } + case e: RejectedExecutionException => + runOk.set(false) + p.setException(e) + null + } - // This is safe: the only thing that can call task.run() is - // executor, the only thing that can raise an interrupt is the - // receiver of this value, which will then be fully initialized. - val javaFuture = try executor.submit(task) - catch { - case e: RejectedExecutionException => - runOk.set(false) - p.setException(e) - null - } + p.setInterruptHandler { + case cause => + if (interruptible || runOk.compareAndSet(true, false)) { + if (p.updateIfEmpty(Throw(cause))) + javaFuture.cancel(true) + } + } - p.setInterruptHandler { - case cause => - if (interruptible || runOk.compareAndSet(true, false)) { - val exc = new CancellationException - exc.initCause(cause) - if (p.updateIfEmpty(Throw(exc))) - javaFuture.cancel(true) - } + p } - - p } override def toString: String = - s"ExecutorServiceFuturePool(interruptible=$interruptible, executor=$executor)" + s"ExecutorServiceFuturePool(interruptible=$interruptible, executor=$executor, forkingScheduler=$forkingScheduler)" override def poolSize: Int = executor match { case tpe: ThreadPoolExecutor => tpe.getPoolSize @@ -190,4 +223,8 @@ class ExecutorServiceFuturePool protected[this] ( case _ => super.numCompletedTasks } + override def numPendingTasks: Long = executor match { + case tpe: ThreadPoolExecutor => tpe.getQueue.size() + case _ => super.numPendingTasks + } } diff --git a/util-core/src/main/scala/com/twitter/util/FutureTransformer.scala b/util-core/src/main/scala/com/twitter/util/FutureTransformer.scala index dcf01653d1..854534f206 100644 --- a/util-core/src/main/scala/com/twitter/util/FutureTransformer.scala +++ b/util-core/src/main/scala/com/twitter/util/FutureTransformer.scala @@ -50,4 +50,3 @@ abstract class FutureTransformer[-A, +B] { */ def handle(throwable: Throwable): B = throw throwable } - diff --git a/util-core/src/main/scala/com/twitter/util/Futures.scala b/util-core/src/main/scala/com/twitter/util/Futures.scala index 741f92f423..67d0199444 100644 --- a/util-core/src/main/scala/com/twitter/util/Futures.scala +++ b/util-core/src/main/scala/com/twitter/util/Futures.scala @@ -1,12 +1,7 @@ package com.twitter.util import java.util.{List => JList, Map => JMap} -import scala.collection.JavaConverters.{ - asScalaBufferConverter, - mapAsJavaMapConverter, - mapAsScalaMapConverter, - seqAsJavaListConverter -} +import scala.jdk.CollectionConverters._ /** * Twitter Future utility methods for ease of use from java @@ -38,9 +33,8 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * underlying futures complete. It fails immediately if any of them * do. */ - def join[A, B](a: Future[A], b: Future[B]): Future[(A, B)] = Future.join(Seq(a, b)).map { _ => - (Await.result(a), Await.result(b)) - } + def join[A, B](a: Future[A], b: Future[B]): Future[(A, B)] = + Future.join(Seq(a, b)).map { _ => (Await.result(a), Await.result(b)) } /** * Join 3 futures. The returned future is complete when all @@ -48,9 +42,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * do. */ def join[A, B, C](a: Future[A], b: Future[B], c: Future[C]): Future[(A, B, C)] = - Future.join(Seq(a, b, c)).map { _ => - (Await.result(a), Await.result(b), Await.result(c)) - } + Future.join(Seq(a, b, c)).map { _ => (Await.result(a), Await.result(b), Await.result(c)) } /** * Join 4 futures. The returned future is complete when all @@ -776,7 +768,7 @@ def join[%s](%s): Future[(%s)] = join(Seq(%s)).map { _ => (%s) }""".format( * to be satisfied and the rest of the futures. */ def select[A](fs: JList[Future[A]]): Future[(Try[A], JList[Future[A]])] = - Future.select(fs.asScala).map { + Future.select(fs.asScala.toSeq).map { case (first, rest) => (first, rest.asJava) } diff --git a/util-core/src/main/scala/com/twitter/util/ImmediateValueFuture.scala b/util-core/src/main/scala/com/twitter/util/ImmediateValueFuture.scala new file mode 100644 index 0000000000..2bd5fe26fb --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/ImmediateValueFuture.scala @@ -0,0 +1,72 @@ +package com.twitter.util + +import scala.util.control.NonFatal + +/** + * Successful Future that contains a value. Transformations are executed immediately (don't go + * through the scheduler). Unlike `Future.const`, Future recursion *will* grow the stack (see + * `ImmediateValueFutureTest` for an example of this) -- it is therefore extremely important that + * you understand the full context of how this Future will be used in order to avoid this. + * + * DO NOT USE THIS without thoroughly understanding the risks! + */ +private[twitter] class ImmediateValueFuture[A](result: A) extends Future[A] { + + private[this] val ReturnResult = Return(result) + + def respond(f: Try[A] => Unit): Future[A] = { + val saved = Local.save() + try { + f(ReturnResult) + } catch Monitor.catcher + finally { + Local.restore(saved) + } + this + } + + override def proxyTo[B >: A](other: Promise[B]): Unit = { + other.update(ReturnResult) + } + + def raise(interrupt: Throwable): Unit = () + + override def rescue[B >: A](rescueException: PartialFunction[Throwable, Future[B]]): Future[B] = { + this + } + + protected def transformTry[B](f: Try[A] => Try[B]): Future[B] = { + val saved = Local.save() + try { + f(ReturnResult) match { + case Return(result) => new ImmediateValueFuture(result) + case t @ Throw(_) => Future.const(t) + } + } catch { + case NonFatal(e) => Future.const(Throw(e)) + } finally { + Local.restore(saved) + } + } + + def transform[B](f: Try[A] => Future[B]): Future[B] = { + val saved = Local.save() + try { + f(ReturnResult) + } catch { + case NonFatal(e) => Future.const(Throw(e)) + } finally { + Local.restore(saved) + } + } + + def poll: Option[Try[A]] = Some(ReturnResult) + + override def toString: String = s"ImmediateValueFuture($result)" + + def ready(timeout: Duration)(implicit permit: Awaitable.CanAwait): this.type = this + + def result(timeout: Duration)(implicit permit: Awaitable.CanAwait): A = result + + def isReady(implicit permit: Awaitable.CanAwait): Boolean = true +} diff --git a/util-core/src/main/scala/com/twitter/util/JavaSingleton.scala b/util-core/src/main/scala/com/twitter/util/JavaSingleton.scala deleted file mode 100644 index 4b7abeef5c..0000000000 --- a/util-core/src/main/scala/com/twitter/util/JavaSingleton.scala +++ /dev/null @@ -1,8 +0,0 @@ -package com.twitter.util - -/** - * A mixin to allow scala objects to be used from java. - */ -trait JavaSingleton { - def get = this -} diff --git a/util-core/src/main/scala/com/twitter/util/LastWriteWinsQueue.scala b/util-core/src/main/scala/com/twitter/util/LastWriteWinsQueue.scala index 13f4d93a89..405f1bc8c8 100644 --- a/util-core/src/main/scala/com/twitter/util/LastWriteWinsQueue.scala +++ b/util-core/src/main/scala/com/twitter/util/LastWriteWinsQueue.scala @@ -7,22 +7,22 @@ import java.util.concurrent.atomic.AtomicReference * When the Queue is full, a push replaces the item. */ class LastWriteWinsQueue[A] extends java.util.Queue[A] { - val item = new AtomicReference[Option[A]](None) + val item: AtomicReference[Option[A]] = new AtomicReference[Option[A]](None) def clear(): Unit = { item.set(None) } - def retainAll(p1: Collection[_]) = throw new UnsupportedOperationException + def retainAll(p1: Collection[_]): Nothing = throw new UnsupportedOperationException - def removeAll(p1: Collection[_]) = throw new UnsupportedOperationException + def removeAll(p1: Collection[_]): Nothing = throw new UnsupportedOperationException - def addAll(p1: Collection[_ <: A]) = throw new UnsupportedOperationException + def addAll(p1: Collection[_ <: A]): Nothing = throw new UnsupportedOperationException - def containsAll(p1: Collection[_]) = + def containsAll(p1: Collection[_]): Boolean = p1.size == 1 && item.get == p1.iterator.next() - def remove(candidate: AnyRef) = { + def remove(candidate: AnyRef): Boolean = { val contained = item.get val containsCandidate = contained.map(_ == candidate).getOrElse(false) if (containsCandidate) { @@ -41,30 +41,30 @@ class LastWriteWinsQueue[A] extends java.util.Queue[A] { } else Array[Any]().asInstanceOf[Array[T with java.lang.Object]] } - def toArray = toArray(new Array[AnyRef](0)) + def toArray: Array[AnyRef with Object] = toArray(new Array[AnyRef](0)) - def iterator = null + def iterator: Null = null - def contains(p1: AnyRef) = false + def contains(p1: AnyRef): Boolean = false - def isEmpty = item.get.isDefined + def isEmpty: Boolean = item.get.isDefined - def size = if (item.get.isDefined) 1 else 0 + def size: Int = if (item.get.isDefined) 1 else 0 - def peek = item.get.getOrElse(null.asInstanceOf[A]) + def peek: A = item.get.getOrElse(null.asInstanceOf[A]) - def element = item.get.getOrElse(throw new NoSuchElementException) + def element: A = item.get.getOrElse(throw new NoSuchElementException) - def poll = item.getAndSet(None).getOrElse(null.asInstanceOf[A]) + def poll: A = item.getAndSet(None).getOrElse(null.asInstanceOf[A]) - def remove = item.getAndSet(None).getOrElse(throw new NoSuchElementException) + def remove: A = item.getAndSet(None).getOrElse(throw new NoSuchElementException) - def offer(p1: A) = { + def offer(p1: A): Boolean = { item.set(Some(p1)) true } - def add(p1: A) = { + def add(p1: A): Boolean = { item.set(Some(p1)) true } diff --git a/util-core/src/main/scala/com/twitter/util/Local.scala b/util-core/src/main/scala/com/twitter/util/Local.scala index 1af139a9c9..98d1f7bd21 100644 --- a/util-core/src/main/scala/com/twitter/util/Local.scala +++ b/util-core/src/main/scala/com/twitter/util/Local.scala @@ -7,7 +7,14 @@ object Local { * when the number of elements is less than 16. * key - Local.Key, value - Option[_] */ - sealed abstract class Context private () { + sealed abstract class Context private ( + final private[util] val resourceTracker: Option[ResourceTracker], + final private[util] val fiber: Fiber) { + private[util] def setResourceTracker(tracker: ResourceTracker): Context + private[util] def removeResourceTracker(): Context + + private[util] def setFiber(f: Fiber): Context + private[util] def get(k: Key): Option[_] private[util] def remove(k: Key): Context private[util] def set(k: Key, v: Some[_]): Context @@ -18,17 +25,42 @@ object Local { /** * The empty Context */ - def empty: Context = EmptyContext + val empty: Context = EmptyContext + + private object EmptyContext extends Context(None, Fiber.Global) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context0(Some(tracker), fiber) + def removeResourceTracker(): Context = new Context0(None, fiber) + + def setFiber(f: Fiber): Context = new Context0(resourceTracker, f) + + def get(k: Key): Option[_] = None + def remove(k: Key): Context = this + def set(k: Key, v: Some[_]): Context = new Context1(resourceTracker, fiber, k, v) + } + + private final class Context0(resourceTracker: Option[ResourceTracker], fiber: Fiber) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context0(Some(tracker), fiber) + def removeResourceTracker(): Context = new Context0(None, fiber) + + def setFiber(f: Fiber): Context = new Context0(resourceTracker, f) - private object EmptyContext extends Context { def get(k: Key): Option[_] = None def remove(k: Key): Context = this - def set(k: Key, v: Some[_]): Context = new Context1(k, v) + def set(k: Key, v: Some[_]): Context = new Context1(resourceTracker, fiber, k, v) } private final class Context1( - k1: Key, v1: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = + new Context1(Some(tracker), fiber, k1, v1) + def removeResourceTracker(): Context = new Context1(None, fiber, k1, v1) + + def setFiber(f: Fiber): Context = new Context1(resourceTracker, f, k1, v1) def get(k: Key): Option[_] = if (k eq k1) v1 else None @@ -37,14 +69,24 @@ object Local { if (k eq k1) EmptyContext else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context1(k1, v) - else new Context2(k1, v1, k, v) + if (k eq k1) new Context1(resourceTracker, fiber, k1, v) + else new Context2(resourceTracker, fiber, k1, v1, k, v) } private final class Context2( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = + new Context2(Some(tracker), fiber, k1, v1, k2, v2) + def removeResourceTracker(): Context = new Context2(None, fiber, k1, v1, k2, v2) + + def setFiber(f: Fiber): Context = + new Context2(resourceTracker, f, k1, v1, k2, v2) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -52,14 +94,14 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context1(k2, v2) - else if (k eq k2) new Context1(k1, v1) + if (k eq k1) new Context1(resourceTracker, fiber, k2, v2) + else if (k eq k2) new Context1(resourceTracker, fiber, k1, v1) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context2(k1, v, k2, v2) - else if (k eq k2) new Context2(k1, v1, k2, v) - else new Context3(k1, v1, k2, v2, k, v) + if (k eq k1) new Context2(resourceTracker, fiber, k1, v, k2, v2) + else if (k eq k2) new Context2(resourceTracker, fiber, k1, v1, k2, v) + else new Context3(resourceTracker, fiber, k1, v1, k2, v2, k, v) } // Script for generating the ContextN @@ -68,8 +110,17 @@ object Local { def mkContext(num: Int): String = { s"""private final class Context$num( + resourceTracker: Option[ResourceTracker], + fiber: Fiber, ${fieldList(num)} - ) extends Context { + ) extends Context(resourceTracker, fiber) { + ${mkSetResourceTracker(num)} + + ${mkRemoveResourceTracker(num)} + + ${mkSetFiber(num)} + + ${mkRemoveFiber(num)} ${mkGet(num)} @@ -87,6 +138,18 @@ object Local { }.mkString(", \n") } + def mkSetResourceTracker(num: Int): String = { + s"def setResourceTracker(tracker: ResourceTracker): Context = new Context$num(Some(tracker), fiber, ${fieldList(num)})" + } + + def mkRemoveResourceTracker(num: Int): String = { + s"def removeResourceTracker(): Context = new Context$num(None, fiber, ${fieldList(num)})" + } + + def mkSetFiber(num: Int): String = { + s"def setFiber(f: Fiber): Context = new Context$num(resourceTracker, f, ${fieldList(num)})" + } + def mkGet(num: Int): String = { val name = s"def get(k: Key): Option[_] = " val content = (1.to(num)).map { i => @@ -100,7 +163,7 @@ object Local { def newCtx(sum: Int, current: Int): String = { (1.to(current - 1) ++: (current + 1).to(sum)).map { i => s"k$i, v$i" - }.mkString("(", ", ", ")") + }.mkString("(resourceTracker, fiber, ", ", ", ")") } val content = (1.to(num)).map { i => s"k$i) new Context${num - 1}${newCtx(num, i)}" @@ -115,7 +178,7 @@ object Local { if (i == current && last) s"k, v" else if (i == current) s"k$i, v" else s"k$i, v$i" - }.mkString("(", ", ", ")") + }.mkString("(resourceTracker, fiber, ", ", ", ")") } val content = (1.to(num)).map { i => s"k$i) new Context$num${newCtx(num, i, false)}" @@ -124,15 +187,42 @@ object Local { } def main(args: Array[String]) { - print(2.to(15).map{ i => mkContext(i)}.mkString) + print(3.to(15).map{ i => mkContext(i)}.mkString) } } */ private final class Context3( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context3( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_]) + + def removeResourceTracker(): Context = + new Context3(None, fiber, k1: Key, v1: Some[_], k2: Key, v2: Some[_], k3: Key, v3: Some[_]) + + def setFiber(f: Fiber): Context = new Context3( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -141,24 +231,65 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context2(k2, v2, k3, v3) - else if (k eq k2) new Context2(k1, v1, k3, v3) - else if (k eq k3) new Context2(k1, v1, k2, v2) + if (k eq k1) new Context2(resourceTracker, fiber, k2, v2, k3, v3) + else if (k eq k2) new Context2(resourceTracker, fiber, k1, v1, k3, v3) + else if (k eq k3) new Context2(resourceTracker, fiber, k1, v1, k2, v2) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context3(k1, v, k2, v2, k3, v3) - else if (k eq k2) new Context3(k1, v1, k2, v, k3, v3) - else if (k eq k3) new Context3(k1, v1, k2, v2, k3, v) - else new Context4(k1, v1, k2, v2, k3, v3, k, v) + if (k eq k1) new Context3(resourceTracker, fiber, k1, v, k2, v2, k3, v3) + else if (k eq k2) new Context3(resourceTracker, fiber, k1, v1, k2, v, k3, v3) + else if (k eq k3) new Context3(resourceTracker, fiber, k1, v1, k2, v2, k3, v) + else new Context4(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k, v) } private final class Context4( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context4( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_]) + + def removeResourceTracker(): Context = new Context4( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_]) + + def setFiber(f: Fiber): Context = new Context4( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -168,27 +299,75 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context3(k2, v2, k3, v3, k4, v4) - else if (k eq k2) new Context3(k1, v1, k3, v3, k4, v4) - else if (k eq k3) new Context3(k1, v1, k2, v2, k4, v4) - else if (k eq k4) new Context3(k1, v1, k2, v2, k3, v3) + if (k eq k1) new Context3(resourceTracker, fiber, k2, v2, k3, v3, k4, v4) + else if (k eq k2) new Context3(resourceTracker, fiber, k1, v1, k3, v3, k4, v4) + else if (k eq k3) new Context3(resourceTracker, fiber, k1, v1, k2, v2, k4, v4) + else if (k eq k4) new Context3(resourceTracker, fiber, k1, v1, k2, v2, k3, v3) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context4(k1, v, k2, v2, k3, v3, k4, v4) - else if (k eq k2) new Context4(k1, v1, k2, v, k3, v3, k4, v4) - else if (k eq k3) new Context4(k1, v1, k2, v2, k3, v, k4, v4) - else if (k eq k4) new Context4(k1, v1, k2, v2, k3, v3, k4, v) - else new Context5(k1, v1, k2, v2, k3, v3, k4, v4, k, v) + if (k eq k1) new Context4(resourceTracker, fiber, k1, v, k2, v2, k3, v3, k4, v4) + else if (k eq k2) new Context4(resourceTracker, fiber, k1, v1, k2, v, k3, v3, k4, v4) + else if (k eq k3) new Context4(resourceTracker, fiber, k1, v1, k2, v2, k3, v, k4, v4) + else if (k eq k4) new Context4(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v) + else new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k, v) } private final class Context5( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context5( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_]) + + def removeResourceTracker(): Context = new Context5( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_]) + + def setFiber(f: Fiber): Context = new Context5( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -199,30 +378,89 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context4(k2, v2, k3, v3, k4, v4, k5, v5) - else if (k eq k2) new Context4(k1, v1, k3, v3, k4, v4, k5, v5) - else if (k eq k3) new Context4(k1, v1, k2, v2, k4, v4, k5, v5) - else if (k eq k4) new Context4(k1, v1, k2, v2, k3, v3, k5, v5) - else if (k eq k5) new Context4(k1, v1, k2, v2, k3, v3, k4, v4) + if (k eq k1) new Context4(resourceTracker, fiber, k2, v2, k3, v3, k4, v4, k5, v5) + else if (k eq k2) new Context4(resourceTracker, fiber, k1, v1, k3, v3, k4, v4, k5, v5) + else if (k eq k3) new Context4(resourceTracker, fiber, k1, v1, k2, v2, k4, v4, k5, v5) + else if (k eq k4) new Context4(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k5, v5) + else if (k eq k5) new Context4(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context5(k1, v, k2, v2, k3, v3, k4, v4, k5, v5) - else if (k eq k2) new Context5(k1, v1, k2, v, k3, v3, k4, v4, k5, v5) - else if (k eq k3) new Context5(k1, v1, k2, v2, k3, v, k4, v4, k5, v5) - else if (k eq k4) new Context5(k1, v1, k2, v2, k3, v3, k4, v, k5, v5) - else if (k eq k5) new Context5(k1, v1, k2, v2, k3, v3, k4, v4, k5, v) - else new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k, v) + if (k eq k1) new Context5(resourceTracker, fiber, k1, v, k2, v2, k3, v3, k4, v4, k5, v5) + else if (k eq k2) + new Context5(resourceTracker, fiber, k1, v1, k2, v, k3, v3, k4, v4, k5, v5) + else if (k eq k3) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v, k4, v4, k5, v5) + else if (k eq k4) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v, k5, v5) + else if (k eq k5) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v) + else new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k, v) } private final class Context6( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context6( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_]) + + def removeResourceTracker(): Context = new Context6( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_]) + + def setFiber(f: Fiber): Context = new Context6( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -234,33 +472,107 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context5(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) - else if (k eq k2) new Context5(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6) - else if (k eq k3) new Context5(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6) - else if (k eq k4) new Context5(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6) - else if (k eq k5) new Context5(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6) - else if (k eq k6) new Context5(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5) + if (k eq k1) new Context5(resourceTracker, fiber, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) + else if (k eq k2) + new Context5(resourceTracker, fiber, k1, v1, k3, v3, k4, v4, k5, v5, k6, v6) + else if (k eq k3) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k4, v4, k5, v5, k6, v6) + else if (k eq k4) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k5, v5, k6, v6) + else if (k eq k5) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k6, v6) + else if (k eq k6) + new Context5(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context6(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) - else if (k eq k2) new Context6(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6) - else if (k eq k3) new Context6(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6) - else if (k eq k4) new Context6(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6) - else if (k eq k5) new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6) - else if (k eq k6) new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v) - else new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k, v) + if (k eq k1) + new Context6(resourceTracker, fiber, k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) + else if (k eq k2) + new Context6(resourceTracker, fiber, k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6) + else if (k eq k3) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6) + else if (k eq k4) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6) + else if (k eq k5) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6) + else if (k eq k6) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v) + else + new Context7(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k, v) } private final class Context7( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context7( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_]) + + def removeResourceTracker(): Context = new Context7( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_]) + + def setFiber(f: Fiber): Context = new Context7( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -273,36 +585,250 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context6(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k2) new Context6(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k3) new Context6(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k4) new Context6(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7) - else if (k eq k5) new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7) - else if (k eq k6) new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7) - else if (k eq k7) new Context6(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) + if (k eq k1) + new Context6(resourceTracker, fiber, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) + else if (k eq k2) + new Context6(resourceTracker, fiber, k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) + else if (k eq k3) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7) + else if (k eq k4) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7) + else if (k eq k5) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7) + else if (k eq k6) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7) + else if (k eq k7) + new Context6(resourceTracker, fiber, k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context7(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k2) new Context7(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k3) new Context7(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7) - else if (k eq k4) new Context7(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7) - else if (k eq k5) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7) - else if (k eq k6) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7) - else if (k eq k7) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v) - else new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k, v) + if (k eq k1) + new Context7( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7) + else if (k eq k2) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7) + else if (k eq k3) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7) + else if (k eq k4) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7) + else if (k eq k5) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7) + else if (k eq k6) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7) + else if (k eq k7) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v) + else + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k, + v) } private final class Context8( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context8( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_]) + + def removeResourceTracker(): Context = new Context8( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_]) + + def setFiber(f: Fiber): Context = new Context8( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -316,39 +842,424 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context7(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k2) new Context7(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k3) new Context7(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k4) new Context7(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k5) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8) - else if (k eq k6) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8) - else if (k eq k7) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8) - else if (k eq k8) new Context7(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7) + if (k eq k1) + new Context7( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k2) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k3) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k4) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k5) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k6) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8) + else if (k eq k7) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8) + else if (k eq k8) + new Context7( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context8(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k2) new Context8(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k3) new Context8(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k4) new Context8(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8) - else if (k eq k5) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8) - else if (k eq k6) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8) - else if (k eq k7) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8) - else if (k eq k8) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v) - else new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k, v) + if (k eq k1) + new Context8( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k2) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k3) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k4) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k5) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8) + else if (k eq k6) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8) + else if (k eq k7) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8) + else if (k eq k8) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v) + else + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k, + v) } private final class Context9( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context9( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_]) + + def removeResourceTracker(): Context = new Context9( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_]) + + def setFiber(f: Fiber): Context = new Context9( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -363,42 +1274,508 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context8(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k2) new Context8(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k3) new Context8(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k4) new Context8(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k5) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k6) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9) - else if (k eq k7) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9) - else if (k eq k8) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9) - else if (k eq k9) new Context8(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8) + if (k eq k1) + new Context8( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k2) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k3) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k4) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k5) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k6) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k7) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9) + else if (k eq k8) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9) + else if (k eq k9) + new Context8( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context9(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k2) new Context9(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k3) new Context9(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k4) new Context9(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k5) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9) - else if (k eq k6) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9) - else if (k eq k7) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9) - else if (k eq k8) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9) - else if (k eq k9) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v) - else new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k, v) + if (k eq k1) + new Context9( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k2) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k3) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k4) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k5) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k6) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9) + else if (k eq k7) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9) + else if (k eq k8) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9) + else if (k eq k9) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v) + else + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k, + v) } private final class Context10( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context10( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_]) + + def removeResourceTracker(): Context = new Context10( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_]) + + def setFiber(f: Fiber): Context = new Context10( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -414,45 +1791,600 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context9(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k2) new Context9(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k3) new Context9(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k4) new Context9(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k5) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k6) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k7) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10) - else if (k eq k8) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10) - else if (k eq k9) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10) - else if (k eq k10) new Context9(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9) + if (k eq k1) + new Context9( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k2) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k3) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k4) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k5) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k6) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k7) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k8) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10) + else if (k eq k9) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10) + else if (k eq k10) + new Context9( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context10(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k2) new Context10(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k3) new Context10(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k4) new Context10(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k5) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k6) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10) - else if (k eq k7) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10) - else if (k eq k8) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10) - else if (k eq k9) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10) - else if (k eq k10) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v) - else new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k, v) + if (k eq k1) + new Context10( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k2) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k3) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k4) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k5) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k6) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k7) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10) + else if (k eq k8) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10) + else if (k eq k9) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10) + else if (k eq k10) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v) + else + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k, + v) } private final class Context11( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_], - k11: Key, v11: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context11( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_]) + + def removeResourceTracker(): Context = new Context11( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_]) + + def setFiber(f: Fiber): Context = new Context11( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -469,48 +2401,700 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context10(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k2) new Context10(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k3) new Context10(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k4) new Context10(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k5) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k6) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k7) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k8) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10, k11, v11) - else if (k eq k9) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10, k11, v11) - else if (k eq k10) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k11, v11) - else if (k eq k11) new Context10(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10) + if (k eq k1) + new Context10( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k2) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k3) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k4) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k5) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k6) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k7) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k8) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k9) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10, + k11, + v11) + else if (k eq k10) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k11, + v11) + else if (k eq k11) + new Context10( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context11(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k2) new Context11(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k3) new Context11(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k4) new Context11(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k5) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k6) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k7) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10, k11, v11) - else if (k eq k8) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10, k11, v11) - else if (k eq k9) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10, k11, v11) - else if (k eq k10) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v, k11, v11) - else if (k eq k11) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v) - else new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k, v) + if (k eq k1) + new Context11( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k2) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k3) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k4) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k5) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k6) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k7) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k8) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10, + k11, + v11) + else if (k eq k9) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10, + k11, + v11) + else if (k eq k10) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v, + k11, + v11) + else if (k eq k11) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v) + else + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k, + v) } private final class Context12( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_], - k11: Key, v11: Some[_], - k12: Key, v12: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context12( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_]) + + def removeResourceTracker(): Context = new Context12( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_]) + + def setFiber(f: Fiber): Context = new Context12( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -528,51 +3112,808 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context11(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k2) new Context11(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k3) new Context11(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k4) new Context11(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k5) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k6) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k7) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k8) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k9) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10, k11, v11, k12, v12) - else if (k eq k10) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k11, v11, k12, v12) - else if (k eq k11) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k12, v12) - else if (k eq k12) new Context11(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11) + if (k eq k1) + new Context11( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k2) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k3) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k4) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k5) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k6) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k7) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k8) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k9) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k10) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k11, + v11, + k12, + v12) + else if (k eq k11) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k12, + v12) + else if (k eq k12) + new Context11( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context12(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k2) new Context12(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k3) new Context12(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k4) new Context12(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k5) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k6) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k7) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k8) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10, k11, v11, k12, v12) - else if (k eq k9) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10, k11, v11, k12, v12) - else if (k eq k10) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v, k11, v11, k12, v12) - else if (k eq k11) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v, k12, v12) - else if (k eq k12) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v) - else new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k, v) + if (k eq k1) + new Context12( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k2) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k3) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k4) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k5) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k6) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k7) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k8) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k9) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10, + k11, + v11, + k12, + v12) + else if (k eq k10) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v, + k11, + v11, + k12, + v12) + else if (k eq k11) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v, + k12, + v12) + else if (k eq k12) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v) + else + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k, + v) } private final class Context13( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_], - k11: Key, v11: Some[_], - k12: Key, v12: Some[_], - k13: Key, v13: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context13( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_]) + + def removeResourceTracker(): Context = new Context13( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_]) + + def setFiber(f: Fiber): Context = new Context13( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -591,54 +3932,924 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context12(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k2) new Context12(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k3) new Context12(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k4) new Context12(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k5) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k6) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k7) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k8) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k9) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k10) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k11, v11, k12, v12, k13, v13) - else if (k eq k11) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k12, v12, k13, v13) - else if (k eq k12) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k13, v13) - else if (k eq k13) new Context12(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12) + if (k eq k1) + new Context12( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k2) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k3) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k4) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k5) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k6) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k7) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k8) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k9) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k10) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k11) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k12, + v12, + k13, + v13) + else if (k eq k12) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k13, + v13) + else if (k eq k13) + new Context12( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context13(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k2) new Context13(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k3) new Context13(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k4) new Context13(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k5) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k6) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k7) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k8) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k9) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10, k11, v11, k12, v12, k13, v13) - else if (k eq k10) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v, k11, v11, k12, v12, k13, v13) - else if (k eq k11) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v, k12, v12, k13, v13) - else if (k eq k12) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v, k13, v13) - else if (k eq k13) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v) - else new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k, v) + if (k eq k1) + new Context13( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k2) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k3) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k4) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k5) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k6) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k7) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k8) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k9) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k10) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v, + k11, + v11, + k12, + v12, + k13, + v13) + else if (k eq k11) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v, + k12, + v12, + k13, + v13) + else if (k eq k12) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v, + k13, + v13) + else if (k eq k13) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v) + else + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k, + v) } private final class Context14( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_], - k11: Key, v11: Some[_], - k12: Key, v12: Some[_], - k13: Key, v13: Some[_], - k14: Key, v14: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context14( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_]) + + def removeResourceTracker(): Context = new Context14( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_]) + + def setFiber(f: Fiber): Context = new Context14( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -658,57 +4869,1048 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context13(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k2) new Context13(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k3) new Context13(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k4) new Context13(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k5) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k6) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k7) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k8) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k9) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k10) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k11) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k12, v12, k13, v13, k14, v14) - else if (k eq k12) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k13, v13, k14, v14) - else if (k eq k13) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k14, v14) - else if (k eq k14) new Context13(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13) + if (k eq k1) + new Context13( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k2) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k3) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k4) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k5) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k6) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k7) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k8) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k9) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k10) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k11) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k12) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k13, + v13, + k14, + v14) + else if (k eq k13) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k14, + v14) + else if (k eq k14) + new Context13( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context14(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k2) new Context14(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k3) new Context14(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k4) new Context14(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k5) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k6) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k7) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k8) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k9) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k10) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v, k11, v11, k12, v12, k13, v13, k14, v14) - else if (k eq k11) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v, k12, v12, k13, v13, k14, v14) - else if (k eq k12) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v, k13, v13, k14, v14) - else if (k eq k13) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v, k14, v14) - else if (k eq k14) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v) - else new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k, v) + if (k eq k1) + new Context14( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k2) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k3) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k4) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k5) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k6) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k7) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k8) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k9) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k10) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k11) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v, + k12, + v12, + k13, + v13, + k14, + v14) + else if (k eq k12) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v, + k13, + v13, + k14, + v14) + else if (k eq k13) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v, + k14, + v14) + else if (k eq k14) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v) + else + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k, + v) } private final class Context15( - k1: Key, v1: Some[_], - k2: Key, v2: Some[_], - k3: Key, v3: Some[_], - k4: Key, v4: Some[_], - k5: Key, v5: Some[_], - k6: Key, v6: Some[_], - k7: Key, v7: Some[_], - k8: Key, v8: Some[_], - k9: Key, v9: Some[_], - k10: Key, v10: Some[_], - k11: Key, v11: Some[_], - k12: Key, v12: Some[_], - k13: Key, v13: Some[_], - k14: Key, v14: Some[_], - k15: Key, v15: Some[_] - ) extends Context { + resourceTracker: Option[ResourceTracker], + fiber: Fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_], + k15: Key, + v15: Some[_]) + extends Context(resourceTracker, fiber) { + def setResourceTracker(tracker: ResourceTracker): Context = new Context15( + Some(tracker), + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_], + k15: Key, + v15: Some[_]) + + def removeResourceTracker(): Context = new Context15( + None, + fiber, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_], + k15: Key, + v15: Some[_]) + + def setFiber(f: Fiber): Context = new Context15( + resourceTracker, + f, + k1: Key, + v1: Some[_], + k2: Key, + v2: Some[_], + k3: Key, + v3: Some[_], + k4: Key, + v4: Some[_], + k5: Key, + v5: Some[_], + k6: Key, + v6: Some[_], + k7: Key, + v7: Some[_], + k8: Key, + v8: Some[_], + k9: Key, + v9: Some[_], + k10: Key, + v10: Some[_], + k11: Key, + v11: Some[_], + k12: Key, + v12: Some[_], + k13: Key, + v13: Some[_], + k14: Key, + v14: Some[_], + k15: Key, + v15: Some[_]) def get(k: Key): Option[_] = if (k eq k1) v1 @@ -729,45 +5931,1010 @@ object Local { else None def remove(k: Key): Context = - if (k eq k1) new Context14(k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k2) new Context14(k1, v1, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k3) new Context14(k1, v1, k2, v2, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k4) new Context14(k1, v1, k2, v2, k3, v3, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k5) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k6) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k7) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k8) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k9) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k10) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k11) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k12) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k13, v13, k14, v14, k15, v15) - else if (k eq k13) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k14, v14, k15, v15) - else if (k eq k14) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k15, v15) - else if (k eq k15) new Context14(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14) + if (k eq k1) + new Context14( + resourceTracker, + fiber, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k2) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k3) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k4) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k5) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k6) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k7) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k8) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k9) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k10) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k11) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k12) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k13) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k14, + v14, + k15, + v15) + else if (k eq k14) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k15, + v15) + else if (k eq k15) + new Context14( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14) else this def set(k: Key, v: Some[_]): Context = - if (k eq k1) new Context15(k1, v, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k2) new Context15(k1, v1, k2, v, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k3) new Context15(k1, v1, k2, v2, k3, v, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k4) new Context15(k1, v1, k2, v2, k3, v3, k4, v, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k5) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k6) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k7) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k8) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k9) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k10) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v, k11, v11, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k11) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v, k12, v12, k13, v13, k14, v14, k15, v15) - else if (k eq k12) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v, k13, v13, k14, v14, k15, v15) - else if (k eq k13) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v, k14, v14, k15, v15) - else if (k eq k14) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v, k15, v15) - else if (k eq k15) new Context15(k1, v1, k2, v2, k3, v3, k4, v4, k5, v5, k6, v6, k7, v7, k8, v8, k9, v9, k10, v10, k11, v11, k12, v12, k13, v13, k14, v14, k15, v) + if (k eq k1) + new Context15( + resourceTracker, + fiber, + k1, + v, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k2) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k3) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k4) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k5) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k6) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k7) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k8) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k9) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k10) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k11) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k12) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v, + k13, + v13, + k14, + v14, + k15, + v15) + else if (k eq k13) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v, + k14, + v14, + k15, + v15) + else if (k eq k14) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v, + k15, + v15) + else if (k eq k15) + new Context15( + resourceTracker, + fiber, + k1, + v1, + k2, + v2, + k3, + v3, + k4, + v4, + k5, + v5, + k6, + v6, + k7, + v7, + k8, + v8, + k9, + v9, + k10, + v10, + k11, + v11, + k12, + v12, + k13, + v13, + k14, + v14, + k15, + v) else new ContextN(k, v, this) } - private final class ContextN( - kN: Key, vN: Some[_], rest: Context - ) extends Context { + private final class ContextN(kN: Key, vN: Some[_], rest: Context) + extends Context(rest.resourceTracker, rest.fiber) { + + def setResourceTracker(tracker: ResourceTracker): Context = + new ContextN(kN, vN, rest.setResourceTracker(tracker)) + def removeResourceTracker(): Context = new ContextN(kN, vN, rest.removeResourceTracker()) + + def setFiber(f: Fiber): Context = new ContextN(kN, vN, rest.setFiber(f)) def get(k: Key): Option[_] = if (k eq kN) vN @@ -788,6 +6955,7 @@ object Local { * Key used in [[Context]], internal. */ private[util] final class Key + private final class ContextRef(var ctx: Context) /** * Represents the current state of all [[Local locals]] for a given @@ -796,44 +6964,73 @@ object Local { * This should be treated as an opaque value and direct modifications * and access are considered verboten. */ - - private[this] val localCtx = new ThreadLocal[Context] { - override def initialValue(): Context = Context.empty + private[this] val localContextRef = new ThreadLocal[ContextRef] { + override def initialValue() = new ContextRef(Context.empty) } /** * Return a snapshot of the current Local state. */ - def save(): Context = localCtx.get + def save(): Context = saveRef().ctx /** * Restore the Local state to a given Context of values. */ - def restore(saved: Context): Unit = localCtx.set(saved) + def restore(saved: Context): Unit = saveRef().ctx = saved + + private def saveRef(): ContextRef = localContextRef.get() private def set(key: Key, value: Some[_]): Unit = { - localCtx.set(localCtx.get.set(key, value)) + val ref = saveRef() + ref.ctx = ref.ctx.set(key, value) } private def get(key: Key): Option[_] = - localCtx.get.get(key) + save().get(key) + + private def clear(key: Key): Unit = { + val ref = saveRef() + ref.ctx = ref.ctx.remove(key) + } - private def clear(key: Key): Unit = - localCtx.set(localCtx.get.remove(key)) + /** + * A lightweight getter scoped to a specific thread. + * + * IMPORTANT: *MUST* be used only by the thread that + * creates it. + */ + final class ThreadLocalGetter[T] private[Local] (key: Key) { + private[this] val owner = Thread.currentThread() + private[this] val ref = saveRef() + private[this] var contextCache = ref.ctx + private[this] var valueCache = contextCache.get(key) + def apply(): Option[T] = { + assert(Thread.currentThread() == owner) + val curr = ref.ctx + if (curr ne contextCache) { + contextCache = curr + valueCache = contextCache.get(key) + } + valueCache.asInstanceOf[Option[T]] + } + } /** * Clear all locals in the current context. */ - def clear(): Unit = localCtx.set(Context.empty) + def clear(): Unit = restore(Context.empty) /** * Execute a block with the given Locals, restoring current values upon completion. */ def let[U](ctx: Context)(f: => U): U = { - val saved = save() - restore(ctx) + val ref = saveRef() + val saved = ref.ctx + ref.ctx = ctx try f - finally restore(saved) + finally { + ref.ctx = saved + } } /** @@ -849,18 +7046,20 @@ object Local { */ def closed[R](fn: () => R): () => R = { val closure = Local.save() - () => - { - val save = Local.save() - Local.restore(closure) - try fn() - finally Local.restore(save) + () => { + val ref = saveRef() + val save = ref.ctx + ref.ctx = closure + try fn() + finally { + ref.ctx = save } + } } } /** - * A Local is a [[ThreadLocal]] whose scope is flexible. The state of all Locals may + * A Local is a `ThreadLocal` whose scope is flexible. The state of all Locals may * be saved or restored onto the current thread by the user. This is useful for * threading Locals through execution contexts. * @@ -900,15 +7099,43 @@ final class Local[T] { */ def apply(): Option[T] = Local.get(key).asInstanceOf[Option[T]] + /** + * Creates a lightweight getter for the current thread. + * + * IMPORTANT: The returned getter *MUST* be used only by + * the thread that creates it. + * + * This method is useful to avoid performance overhead if + * a local needs to be accessed several times by the same + * thread. + */ + def threadLocalGetter(): Local.ThreadLocalGetter[T] = + new Local.ThreadLocalGetter(key) + /** * Execute a block with a specific Local value, restoring the current state * upon completion. */ def let[U](value: T)(f: => U): U = { - val saved = apply() - update(value) + val ref = Local.saveRef() + val oldCtx = ref.ctx + val newCtx = oldCtx.set(key, Some(value)) + ref.ctx = newCtx try f - finally set(saved) + finally { + val now = ref.ctx + if (newCtx eq now) { + // Fast path: no other ctx modifications to worry about + ref.ctx = oldCtx + } else { + // Another element was updated in the meantime. We just have to set the old value. + val next = oldCtx.get(key) match { + case s @ Some(_) => now.set(key, s) + case None => now.remove(key) + } + ref.ctx = next + } + } } /** @@ -916,10 +7143,25 @@ final class Local[T] { * completion. */ def letClear[U](f: => U): U = { - val saved = apply() - clear() + val ref = Local.saveRef() + val oldCtx = ref.ctx + val newCtx = oldCtx.remove(key) + ref.ctx = newCtx try f - finally set(saved) + finally { + val now = ref.ctx + if (newCtx eq now) { + // Fast path: no other ctx modifications to worry about + ref.ctx = oldCtx + } else { + // Another element was updated in the meantime. We just have to set the old value. + val next = oldCtx.get(key) match { + case s @ Some(_) => now.set(key, s) + case None => now.remove(key) + } + ref.ctx = next + } + } } /** diff --git a/util-core/src/main/scala/com/twitter/util/Memoize.scala b/util-core/src/main/scala/com/twitter/util/Memoize.scala index 66d7b49d48..97e9a09bfa 100644 --- a/util-core/src/main/scala/com/twitter/util/Memoize.scala +++ b/util-core/src/main/scala/com/twitter/util/Memoize.scala @@ -36,6 +36,12 @@ object Memoize { * [[scala.collection.immutable.Map]] */ def snap: Map[A, B] + + /** + * Returns the most recent size of the currently memoized computations. + * Useful for stats reporting. + */ + def size: Long } /** @@ -79,6 +85,8 @@ object Memoize { new Snappable[A, B] { private[this] var memo = Map.empty[A, Either[GuardedCountDownLatch, B]] + override def size: Long = memo.size + def snap: Map[A, B] = synchronized(memo) collect { case (a, Right(b)) => (a, b) @@ -156,4 +164,21 @@ object Memoize { case _ => missing(a) } } + + /** + * Thread-safe memoization for a function taking Class objects, + * using [[java.lang.ClassValue]]. + * + * @param f the function to memoize using a `ClassValue` + * @tparam T the return type of the function + * @return a memoized function over classes that stores memoized values + * in a `ClassValue` + */ + def classValue[T](f: Class[_] => T): Class[_] => T = { + val cv = new ClassValue[T] { + def computeValue(c: Class[_]): T = f(c) + } + + cv.get + } } diff --git a/util-core/src/main/scala/com/twitter/util/Monitor.scala b/util-core/src/main/scala/com/twitter/util/Monitor.scala index ad8e02251f..f708cf515f 100644 --- a/util-core/src/main/scala/com/twitter/util/Monitor.scala +++ b/util-core/src/main/scala/com/twitter/util/Monitor.scala @@ -17,7 +17,7 @@ case class MonitorException(handlingExc: Throwable, monitorExc: Throwable) /** * A Monitor is a composable exception handler. It is independent of * position, divorced from the notion of a call stack. Monitors do - * not recover values from a failed computations: It handles only true + * not recover values from failed computations: It handles only true * exceptions that may require cleanup. * * @see [[AbstractMonitor]] for an API friendly to creating instances from Java. @@ -31,7 +31,7 @@ trait Monitor { self => */ def handle(exc: Throwable): Boolean - private[this] val someSelf = Some(self) + private[this] val someSelf: Some[Monitor] = Some(self) /** * Run `f` inside of the monitor context. If `f` throws diff --git a/util-core/src/main/scala/com/twitter/util/NetUtil.scala b/util-core/src/main/scala/com/twitter/util/NetUtil.scala index 4c87bd7710..86de343726 100644 --- a/util-core/src/main/scala/com/twitter/util/NetUtil.scala +++ b/util-core/src/main/scala/com/twitter/util/NetUtil.scala @@ -131,9 +131,7 @@ object NetUtil { isIpInBlock(inetAddressToInt(inetAddress), ipBlock) def isIpInBlocks(ip: Int, ipBlocks: Iterable[(Int, Int)]): Boolean = { - ipBlocks exists { ipBlock => - isIpInBlock(ip, ipBlock) - } + ipBlocks exists { ipBlock => isIpInBlock(ip, ipBlock) } } def isIpInBlocks(ip: String, ipBlocks: Iterable[(Int, Int)]): Boolean = { diff --git a/util-core/src/main/scala/com/twitter/util/NoFuture.scala b/util-core/src/main/scala/com/twitter/util/NoFuture.scala index c498552ca5..5cf3dfba81 100644 --- a/util-core/src/main/scala/com/twitter/util/NoFuture.scala +++ b/util-core/src/main/scala/com/twitter/util/NoFuture.scala @@ -9,6 +9,7 @@ package com.twitter.util class NoFuture extends Future[Nothing] { def respond(k: Try[Nothing] => Unit): Future[Nothing] = this def transform[B](f: Try[Nothing] => Future[B]): Future[B] = this + protected def transformTry[B](f: Try[Nothing] => Try[B]): Future[B] = this def raise(interrupt: Throwable): Unit = () diff --git a/util-core/src/main/scala/com/twitter/util/Pool.scala b/util-core/src/main/scala/com/twitter/util/Pool.scala index 02659d0eb0..7ee7ec87bb 100644 --- a/util-core/src/main/scala/com/twitter/util/Pool.scala +++ b/util-core/src/main/scala/com/twitter/util/Pool.scala @@ -1,6 +1,7 @@ package com.twitter.util import scala.collection.mutable +import scala.collection.Iterable trait Pool[A] { def reserve(): Future[A] @@ -8,7 +9,7 @@ trait Pool[A] { } class SimplePool[A](items: mutable.Queue[Future[A]]) extends Pool[A] { - def this(initialItems: Seq[A]) = this { + def this(initialItems: Iterable[A]) = this { val queue = new mutable.Queue[Future[A]] queue ++= initialItems.map(Future(_)) queue @@ -41,7 +42,7 @@ class SimplePool[A](items: mutable.Queue[Future[A]]) extends Pool[A] { } abstract class FactoryPool[A](numItems: Int) extends Pool[A] { - private val healthyQueue = new HealthyQueue[A](makeItem, numItems, isHealthy) + private val healthyQueue = new HealthyQueue[A](makeItem _, numItems, isHealthy) private val simplePool = new SimplePool[A](healthyQueue) def reserve(): Future[A] = simplePool.reserve() @@ -53,39 +54,3 @@ abstract class FactoryPool[A](numItems: Int) extends Pool[A] { protected def makeItem(): Future[A] protected def isHealthy(a: A): Boolean } - -private class HealthyQueue[A](makeItem: () => Future[A], numItems: Int, isHealthy: A => Boolean) - extends mutable.Queue[Future[A]] { - - 0.until(numItems) foreach { _ => - this += makeItem() - } - - override def +=(elem: Future[A]): HealthyQueue.this.type = synchronized { - super.+=(elem) - } - - override def +=:(elem: Future[A]): HealthyQueue.this.type = synchronized { - super.+=:(elem) - } - - override def enqueue(elems: Future[A]*): Unit = synchronized { - super.enqueue(elems: _*) - } - - override def dequeue(): Future[A] = synchronized { - if (isEmpty) throw new NoSuchElementException("queue empty") - - super.dequeue() flatMap { item => - if (isHealthy(item)) { - Future(item) - } else { - val item = makeItem() - synchronized { - enqueue(item) - dequeue() - } - } - } - } -} diff --git a/util-core/src/main/scala/com/twitter/util/Promise.scala b/util-core/src/main/scala/com/twitter/util/Promise.scala index 6bb12e5af3..bb51047232 100644 --- a/util-core/src/main/scala/com/twitter/util/Promise.scala +++ b/util-core/src/main/scala/com/twitter/util/Promise.scala @@ -1,17 +1,26 @@ package com.twitter.util -import com.twitter.concurrent.Scheduler import scala.annotation.tailrec import scala.runtime.NonLocalReturnControl import scala.util.control.NonFatal +/** + * @define PromiseInterruptsSingleArgLink + * [[com.twitter.util.Promise.interrupts[A](f:com\.twitter\.util\.Future[_]):com\.twitter\.util\.Promise[A]* interrupts(Future)]] + * + * @define PromiseInterruptsTwoArgsLink + * [[com.twitter.util.Promise.interrupts[A](a:com\.twitter\.util\.Future[_],b:com\.twitter\.util\.Future[_]):com\.twitter\.util\.Promise[A]* interrupts(Future,Future)]] + * + * @define PromiseInterruptsVarargsLink + * [[com.twitter.util.Promise.interrupts[A](fs:com\.twitter\.util\.Future[_]*):com\.twitter\.util\.Promise[A]* interrupts(Future*)]] + */ object Promise { /** * Embeds an "interrupt handler" into a [[Promise]]. * - * This is a total handler such that it's defined on any [[Throwable]]. Use - * [[Promise.setInterruptHandler()]] if you need to leave an interrupt handler + * This is a total handler such that it's defined on any `Throwable`. Use + * [[Promise.setInterruptHandler]] if you need to leave an interrupt handler * undefined for certain types of exceptions. * * Example: (`p` and `q` are equivalent, but `p` allocates less): @@ -54,7 +63,7 @@ object Promise { /** * A persistent queue of continuations (i.e., `K`). */ - private[util] sealed trait WaitQueue[-A] { + private[util] sealed abstract class WaitQueue[-A] { def first: K[A] def rest: WaitQueue[A] @@ -82,22 +91,26 @@ object Promise { loop(this, WaitQueue.empty) } - final def runInScheduler(t: Try[A]): Unit = - Scheduler.submit(new Runnable() { def run(): Unit = WaitQueue.this.run(t) }) - @tailrec - private def run(t: Try[A]): Unit = + final def runInScheduler(t: Try[A]): Unit = { if (this ne WaitQueue.Empty) { - first(t) - rest.run(t) + this.first.saved.fiber.submitTask(() => first(t)) + rest.runInScheduler(t) } + } final override def toString: String = s"WaitQueue(size=$size)" } private[util] object WaitQueue { - val Empty: WaitQueue[Nothing] = new WaitQueue[Nothing] { + // Works the best when length(from) < length(to) + @tailrec + final def merge[A](from: WaitQueue[A], to: WaitQueue[A]): WaitQueue[A] = + if (from eq Empty) to + else merge(from.rest, WaitQueue(from.first, to)) + + object Empty extends WaitQueue[Nothing] { final def first: K[Nothing] = throw new IllegalStateException("WaitQueue.Empty") @@ -124,16 +137,17 @@ object Promise { * it will make "Linked" and "Waiting" state cases ambiguous. This, however, * may change following the further performance improvements. */ - private[util] trait K[-A] extends (Try[A] => Unit) with WaitQueue[A] { + private[util] abstract class K[-A](val saved: Local.Context) extends WaitQueue[A] { final def first: K[A] = this final def rest: WaitQueue[A] = WaitQueue.empty + def apply(r: Try[A]): Unit } /** * A template trait for [[com.twitter.util.Promise Promises]] that are derived * and capable of being detached from other Promises. */ - trait Detachable { _: Promise[_] => + trait Detachable { self: Promise[_] => /** * Returns true if successfully detached, will return true at most once. @@ -153,7 +167,7 @@ object Promise { // It's not possible (yet) to embed K[A] into Promise because // Promise[A] (Linked) and WaitQueue (Waiting) states become ambiguous. - private[this] val k = new K[A] { + private[this] val k = new K[A](Local.Context.empty) { // This is only called after the underlying has been successfully satisfied def apply(result: Try[A]): Unit = self.update(result) } @@ -172,16 +186,17 @@ object Promise { with Detachable with (Try[A] => Unit) { - private[this] var detached: Boolean = false + // 0 represents not yet detached, 1 represents detached. + @volatile + private[this] var alreadyDetached: Int = 0 - def detach(): Boolean = synchronized { - if (detached) { - false - } else { - detached = true - true - } - } + // We add this method so that the Scala 3 compiler doesn't optimize + // away the `alreadyDetached` field which isn't explictly used within + // the scope of this class + private[this] def detached: Boolean = alreadyDetached == 1 + + def detach(): Boolean = + unsafe.compareAndSwapInt(this, detachedFutureOffset, 0, 1) def apply(result: Try[A]): Unit = if (detach()) update(result) @@ -196,13 +211,29 @@ object Promise { * @param k the closure to invoke in the saved context, with the * provided result */ - private class Monitored[A](saved: Local.Context, k: Try[A] => Unit) extends K[A] { + private final class Monitored[A](saved: Local.Context, k: Try[A] => Unit) extends K[A](saved) { + def apply(result: Try[A]): Unit = { - val current = Local.save() - Local.restore(saved) - try k(result) - catch Monitor.catcher - finally Local.restore(current) + Local.let(saved) { + try k(result) + catch Monitor.catcher + } + } + } + + private abstract class Transformer[A](saved: Local.Context) extends K[A](saved) { + + protected[this] def k(r: Try[A]): Unit + + final def apply(result: Try[A]): Unit = { + Local.let(saved) { + try k(result) + catch { + case t: Throwable => + Monitor.handle(t) + throw t + } + } } } @@ -210,16 +241,17 @@ object Promise { * A transforming continuation. * * @param saved The saved local context of the invocation site - * @param promise The Promise for the transformed value * @param f The closure to invoke to produce the Future of the transformed value. + * @param promise The Promise for the transformed value */ - private class Transformer[A, B]( + private final class FutureTransformer[A, B]( saved: Local.Context, - promise: Promise[B], - f: Try[A] => Future[B] - ) extends K[A] { + f: Try[A] => Future[B], + promise: Promise[B]) + extends Transformer[A](saved) { - private[this] def k(r: Try[A]): Unit = { + protected[this] def k(r: Try[A]): Unit = + // The promise can be fulfilled only by the transformer, so it's safe to use `become` here promise.become( try f(r) catch { @@ -227,29 +259,26 @@ object Promise { case NonFatal(e) => Future.exception(e) } ) - } + } - def apply(result: Try[A]): Unit = { - val current = Local.save() - Local.restore(saved) - try k(result) - catch { - case t: Throwable => - Monitor.handle(t) - throw t - } finally Local.restore(current) + private final class TryTransformer[A, B]( + saved: Local.Context, + f: Try[A] => Try[B], + promise: Promise[B]) + extends Transformer[A](saved) { + + protected[this] def k(r: Try[A]): Unit = { + // The promise can be fulfilled only by the transformer, so it's safe to use `update` here + promise.update( + try f(r) + catch { + case e: NonLocalReturnControl[_] => Throw(new FutureNonLocalReturnControl(e)) + case NonFatal(e) => Throw(e) + } + ) } } - /* - * Implementation notes: - * - * While these various states for `Promises.state` should be nicely modeled - * as a sealed type, by omitting the wrappers around `Done` and `Linked` - * we are able to save significant number of allocations for slightly - * more unreadable code localized within `Promise`. - */ - /** * An unsatisfied [[Promise]] which has an interrupt handler attached to it. * `waitq` represents the continuations that should be run once it @@ -257,8 +286,8 @@ object Promise { */ private class Interruptible[A]( val waitq: WaitQueue[A], - val handler: PartialFunction[Throwable, Unit] - ) + val handler: PartialFunction[Throwable, Unit], + val saved: Local.Context) /** * An unsatisfied [[Promise]] which forwards interrupts to `other`. @@ -277,38 +306,11 @@ object Promise { private val unsafe: sun.misc.Unsafe = Unsafe() private val stateOff: Long = unsafe.objectFieldOffset(classOf[Promise[_]].getDeclaredField("state")) - private val AlwaysUnit: Any => Unit = _ => () - - sealed trait Responder[A] { this: Future[A] => - protected final def continueAll(wq: WaitQueue[A]): Unit = { - var ks = wq - while (ks ne WaitQueue.Empty) { - continue(ks.first) - ks = ks.rest - } - } - - protected def continue(k: K[A]): Unit - - /** - * Note: exceptions in responds are monitored. That is, if the - * computation `k` throws a raw (ie. not encoded in a Future) - * exception, it is handled by the current monitor, see - * [[Monitor]] for details. - */ - def respond(k: Try[A] => Unit): Future[A] = { - continue(new Monitored(Local.save(), k)) - this - } - - def transform[B](f: Try[A] => Future[B]): Future[B] = { - val promise = interrupts[B](this) - continue(new Transformer(Local.save(), promise, f)) + private val detachedFutureOffset: Long = + unsafe.objectFieldOffset(classOf[DetachableFuture[_]].getDeclaredField("alreadyDetached")) - promise - } - } + private val AlwaysUnit: Any => Unit = _ => () // PUBLIC API @@ -324,19 +326,17 @@ object Promise { /** * Single-arg version to avoid object creation and take advantage of `forwardInterruptsTo`. * - * @see [[interrupts(Future, Future)]] - * @see [[interrupts(Future*)]] + * @see $PromiseInterruptsTwoArgsLink + * @see $PromiseInterruptsVarargsLink */ - def interrupts[A](f: Future[_]): Promise[A] = new Promise[A] { - forwardInterruptsTo(f) - } + def interrupts[A](f: Future[_]): Promise[A] = new Promise[A](f) /** * Create a promise that interrupts `a` and `b` futures. In particular: * the returned promise handles an interrupt when either `a` or `b` does. * - * @see [[interrupts(Future)]] - * @see [[interrupts(Future*)]] + * @see $PromiseInterruptsSingleArgLink + * @see $PromiseInterruptsVarargsLink */ def interrupts[A](a: Future[_], b: Future[_]): Promise[A] = new Promise[A] with InterruptHandler { @@ -350,8 +350,8 @@ object Promise { * Create a promise that interrupts all of `fs`. In particular: * the returned promise handles an interrupt when any of `fs` do. * - * @see [[interrupts(Future)]] - * @see [[interrupts(Future, Future)]] + * @see $PromiseInterruptsSingleArgLink + * @see $PromiseInterruptsTwoArgsLink */ def interrupts[A](fs: Future[_]*): Promise[A] = new Promise[A] with InterruptHandler { @@ -404,9 +404,9 @@ object Promise { * `Interrupted`, `Transforming`, `Done` and `Linked` where `Interruptible`, * `Interrupted`, and `Transforming` are variants of `Waiting` to deal with future * interrupts. Promises are concurrency-safe, using lock-free operations - * throughout. Callback dispatch is scheduled with [[Scheduler]]. + * throughout. Callback dispatch is scheduled with [[com.twitter.concurrent.Scheduler]]. * - * Waiters (i.e., continuations) are stored in a [[Promise.WaitQueue]] and + * Waiters (i.e., continuations) are stored in a `Promise.WaitQueue` and * executed in the LIFO order. * * `Promise.become` merges two promises: they are declared equivalent. @@ -415,37 +415,87 @@ object Promise { * space leak is incurred from `flatMap` in the tail position since * intermediate promises are merged into the root promise. */ -class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[A]] { +class Promise[A] extends Future[A] with Updatable[Try[A]] { import Promise._ + import ResourceTracker.wrapAndMeasureUsage - // Note: this will always be one of: - // - WaitQueue (Waiting) - // - Interrupted - // - Interruptible - // - Transforming - // - Try[A] (Done) - // - Promise[A] + /** + * Note: exceptions in responds are monitored. That is, if the + * computation `k` throws a raw (ie. not encoded in a Future) + * exception, it is handled by the current monitor, see + * [[Monitor]] for details. + */ + final def respond(k: Try[A] => Unit): Future[A] = { + val saved = Local.save() + val tracker = saved.resourceTracker + if (tracker eq None) continue(new Monitored(saved, k)) + else continue(new Monitored(saved, wrapAndMeasureUsage(k, tracker.get))) + this + } + + final def transform[B](f: Try[A] => Future[B]): Future[B] = { + val promise = interrupts[B](this) + val saved = Local.save() + val tracker = saved.resourceTracker + if (tracker eq None) continue(new FutureTransformer(saved, f, promise)) + else continue(new FutureTransformer(saved, wrapAndMeasureUsage(f, tracker.get), promise)) + promise + } + + final protected def transformTry[B](f: Try[A] => Try[B]): Future[B] = { + val promise = interrupts[B](this) + val saved = Local.save() + val tracker = saved.resourceTracker + if (tracker eq None) continue(new TryTransformer(saved, f, promise)) + else continue(new TryTransformer(saved, wrapAndMeasureUsage(f, tracker.get), promise)) + promise + } + + // Promise state encoding notes: + // + // There are six (6) possible promise states that we encode as raw objects (no state class + // wrappers) for the sake of reducing allocations. This is why, for example, state "Linked" + // is encoded as directly 'Promise[_]' and not as 'case class Linked(to: Promise[_])'. + // + // When we transition between states, we take into consideration the order in which we match + // against the possible state cases. To determine the most optimal match oder, we instrumented + // the Promise code to keep track of the histogram of the most frequent states. + // + // Here is the data we've collected from one of our production services: + // - Transforming: 73% + // - Waiting: 13% + // - Done: 8% + // - Interruptible: 4% + // - Linked: 1% + // - Interrupted: <1% + // + // We use this exact match order (from most frequent to less frequent) in the Promise code. @volatile private[this] var state: Any = WaitQueue.empty[A] private def theState(): Any = state - def this(handleInterrupt: PartialFunction[Throwable, Unit]) { + private[util] def this(forwardInterrupts: Future[_]) = { + this() + this.state = new Transforming[A](WaitQueue.empty, forwardInterrupts) + } + + def this(handleInterrupt: PartialFunction[Throwable, Unit]) = { this() - this.state = new Interruptible[A](WaitQueue.empty, handleInterrupt) + this.state = new Interruptible[A](WaitQueue.empty, handleInterrupt, Local.save()) } - def this(result: Try[A]) { + def this(result: Try[A]) = { this() this.state = result } override def toString: String = { val theState = state match { - case waitq: WaitQueue[A] => s"Waiting($waitq)" - case s: Interruptible[A] => s"Interruptible(${s.waitq},${s.handler})" case s: Transforming[A] => s"Transforming(${s.waitq},${s.other})" - case s: Interrupted[A] => s"Interrupted(${s.waitq},${s.signal})" + case waitq: Promise.WaitQueue[A] => s"Waiting($waitq)" case res: Try[A] => s"Done($res)" + case s: Interruptible[A] => s"Interruptible(${s.waitq},${s.handler})" case p: Promise[A] => s"Linked(${p.toString})" + case s: Interrupted[A] => s"Interrupted(${s.waitq},${s.signal})" } s"Promise@$hashCode(state=$theState)" } @@ -459,36 +509,46 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ * * @param f the new interrupt handler */ + final def setInterruptHandler(f: PartialFunction[Throwable, Unit]): Unit = + setInterruptHandler0(f, WaitQueue.empty) + @tailrec - final def setInterruptHandler(f: PartialFunction[Throwable, Unit]): Unit = state match { - case waitq: WaitQueue[A] => - if (!cas(waitq, new Interruptible(waitq, f))) - setInterruptHandler(f) + private final def setInterruptHandler0( + f: PartialFunction[Throwable, Unit], + wq: WaitQueue[A] + ): Unit = + state match { + case s: Transforming[A] => + if (!cas(s, new Interruptible(WaitQueue.merge(wq, s.waitq), f, Local.save()))) + setInterruptHandler0(f, wq) - case s: Interruptible[A] => - if (!cas(s, new Interruptible(s.waitq, f))) - setInterruptHandler(f) + case waitq: WaitQueue[A] => + if (!cas(waitq, new Interruptible(WaitQueue.merge(wq, waitq), f, Local.save()))) + setInterruptHandler0(f, wq) - case s: Transforming[A] => - if (!cas(s, new Interruptible(s.waitq, f))) - setInterruptHandler(f) + case t: Try[A] /* Done */ => wq.runInScheduler(t) - case s: Interrupted[A] => - f.applyOrElse(s.signal, Promise.AlwaysUnit) + case s: Interruptible[A] => + if (!cas(s, new Interruptible(WaitQueue.merge(wq, s.waitq), f, Local.save()))) + setInterruptHandler0(f, wq) - case _: Try[A] /* Done */ => // ignore + case p: Promise[A] /* Linked */ => p.setInterruptHandler0(f, wq) - case p: Promise[A] /* Linked */ => p.setInterruptHandler(f) - } + case s: Interrupted[A] => + if ((wq ne WaitQueue.Empty) && + !cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), s.signal))) + setInterruptHandler0(f, wq) + else f.applyOrElse(s.signal, Promise.AlwaysUnit) + } // Useful for debugging waitq. private[util] def waitqLength: Int = state match { - case waitq: WaitQueue[A] => waitq.size - case s: Interruptible[A] => s.waitq.size case s: Transforming[A] => s.waitq.size - case s: Interrupted[A] => s.waitq.size + case waitq: WaitQueue[A] => waitq.size case _: Try[A] /* Done */ => 0 + case s: Interruptible[A] => s.waitq.size case _: Promise[A] /* Linked */ => 0 + case s: Interrupted[A] => s.waitq.size } /** @@ -499,92 +559,106 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ * * @param other the Future to which interrupts are forwarded. */ - @tailrec final def forwardInterruptsTo(other: Future[_]): Unit = { - // This reduces allocations in the common case. - if (other.isDefined) return - state match { - case waitq: WaitQueue[A] => - if (!cas(waitq, new Transforming(waitq, other))) - forwardInterruptsTo(other) - - case s: Interruptible[A] => - if (!cas(s, new Transforming(s.waitq, other))) - forwardInterruptsTo(other) - - case s: Transforming[A] => - if (!cas(s, new Transforming(s.waitq, other))) - forwardInterruptsTo(other) + final def forwardInterruptsTo(other: Future[_]): Unit = + forwardInterruptsTo0(other, WaitQueue.empty) - case s: Interrupted[_] => - other.raise(s.signal) + @tailrec private final def forwardInterruptsTo0(other: Future[_], wq: WaitQueue[A]): Unit = { + // This reduces allocations in the common case. + if (other.isDefined) continue(wq) + else + state match { + case s: Transforming[A] => + if (!cas(s, new Transforming(WaitQueue.merge(wq, s.waitq), other))) + forwardInterruptsTo0(other, wq) + + case waitq: WaitQueue[A] => + if (!cas(waitq, new Transforming(WaitQueue.merge(wq, waitq), other))) + forwardInterruptsTo0(other, wq) + + case t: Try[A] /* Done */ => wq.runInScheduler(t) + + case s: Interruptible[A] => + if (!cas(s, new Transforming(WaitQueue.merge(wq, s.waitq), other))) + forwardInterruptsTo0(other, wq) + + case p: Promise[A] /* Linked */ => p.forwardInterruptsTo0(other, wq) + + case s: Interrupted[_] => + if ((wq ne WaitQueue.Empty) && + !cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), s.signal))) + forwardInterruptsTo0(other, wq) + else other.raise(s.signal) + } + } - case _: Try[A] /* Done */ => () // ignore + final def raise(intr: Throwable): Unit = raise0(intr, WaitQueue.empty) - case p: Promise[A] /* Linked */ => p.forwardInterruptsTo(other) - } - } + // This is implemented as a private method on Promise so that the Scala 3 + // compiler recognizes this as a tail-recursive implementation. `raise` is + // a definition on `Future` instead of `Promise`, which the Scala 3 + // compiler errors on with @tailrec annocation on the `Transforming` case, + // where `s.other.raise` is invoked on a Future and not Promise. See CSL-11192. + @tailrec private final def raise0(intr: Throwable, wq: WaitQueue[A]): Unit = state match { + case s: Transforming[A] => + if (!cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), intr))) raise0(intr, wq) + else s.other.raise(intr) - @tailrec final def raise(intr: Throwable): Unit = state match { case waitq: WaitQueue[A] => - if (!cas(waitq, new Interrupted(waitq, intr))) - raise(intr) + if (!cas(waitq, new Interrupted(WaitQueue.merge(wq, waitq), intr))) + raise0(intr, wq) + + case t: Try[A] /* Done */ => wq.runInScheduler(t) case s: Interruptible[A] => - if (!cas(s, new Interrupted(s.waitq, intr))) raise(intr) + if (!cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), intr))) raise0(intr, wq) else { - s.handler.applyOrElse(intr, Promise.AlwaysUnit) + Local.let(s.saved) { + s.handler.applyOrElse(intr, Promise.AlwaysUnit) + } } - case s: Transforming[A] => - if (!cas(s, new Interrupted(s.waitq, intr))) raise(intr) - else { - s.other.raise(intr) - } + case p: Promise[A] /* Linked */ => p.raise0(intr, wq) case s: Interrupted[A] => - if (!cas(s, new Interrupted(s.waitq, intr))) - raise(intr) - - case _: Try[A] /* Done */ => () // nothing to do, as its already satisfied. - - case p: Promise[A] /* Linked */ => p.raise(intr) + if (!cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), intr))) + raise0(intr, wq) } @tailrec protected[Promise] final def detach(k: K[A]): Boolean = state match { + case s: Transforming[A] => + if (!cas(s, new Transforming(s.waitq.remove(k), s.other))) + detach(k) + else + s.waitq.contains(k) + case waitq: WaitQueue[A] => if (!cas(waitq, waitq.remove(k))) detach(k) else waitq.contains(k) + case _: Try[A] /* Done */ => false + case s: Interruptible[A] => - if (!cas(s, new Interruptible(s.waitq.remove(k), s.handler))) + if (!cas(s, new Interruptible(s.waitq.remove(k), s.handler, s.saved))) detach(k) else s.waitq.contains(k) - case s: Transforming[A] => - if (!cas(s, new Transforming(s.waitq.remove(k), s.other))) - detach(k) - else - s.waitq.contains(k) + case p: Promise[A] /* Linked */ => p.detach(k) case s: Interrupted[A] => if (!cas(s, new Interrupted(s.waitq.remove(k), s.signal))) detach(k) else s.waitq.contains(k) - - case _: Try[A] /* Done */ => false - - case p: Promise[A] /* Linked */ => p.detach(k) } // Awaitable @throws(classOf[TimeoutException]) @throws(classOf[InterruptedException]) def ready(timeout: Duration)(implicit permit: Awaitable.CanAwait): this.type = state match { - case _: WaitQueue[A] | _: Interruptible[A] | _: Interrupted[A] | _: Transforming[A] => + case _: Transforming[A] | _: WaitQueue[A] | _: Interruptible[A] | _: Interrupted[A] => val condition = new ReleaseOnApplyCDL[A] respond(condition) @@ -595,12 +669,12 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ // Await.result(Future.Done.map(Predef.identity)) // } // - Scheduler.flush() + Fiber.flush() if (condition.await(timeout.inNanoseconds, java.util.concurrent.TimeUnit.NANOSECONDS)) this else throw new TimeoutException(timeout.toString) - case res: Try[A] /* Done */ => this + case _: Try[A] /* Done */ => this case p: Promise[A] /* Linked */ => p.ready(timeout); this } @@ -618,10 +692,10 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ * Returns this promise's interrupt if it is interrupted. */ def isInterrupted: Option[Throwable] = state match { - case _: WaitQueue[A] | _: Interruptible[A] | _: Transforming[A] => None - case s: Interrupted[A] => Some(s.signal) + case _: Transforming[A] | _: WaitQueue[A] | _: Interruptible[A] => None case _: Try[A] /* Done */ => None case p: Promise[A] /* Linked */ => p.isInterrupted + case s: Interrupted[A] => Some(s.signal) } /** @@ -667,14 +741,16 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ * @see [[com.twitter.util.Future.proxyTo]] */ def become(other: Future[A]): Unit = { - if (isDefined) { - val current = Await.result(liftToTry) - throw new IllegalStateException(s"cannot become() on an already satisfied promise: $current") - } if (other.isInstanceOf[Promise[_]]) { + if (isDefined) { + val current = Await.result(liftToTry) + throw new IllegalStateException( + s"cannot become() on an already satisfied promise: $current") + } val that = other.asInstanceOf[Promise[A]] that.link(compress()) } else { + // avoid an extra call to `isDefined` as `proxyTo` checks other.proxyTo(this) forwardInterruptsTo(other) } @@ -683,14 +759,14 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ /** * Populate the Promise with the given result. * - * @throws ImmutableResult if the Promise is already populated + * @throws Promise.ImmutableResult if the Promise is already populated */ def setValue(result: A): Unit = update(Return(result)) /** * Populate the Promise with the given exception. * - * @throws ImmutableResult if the Promise is already populated + * @throws Promise.ImmutableResult if the Promise is already populated */ def setException(throwable: Throwable): Unit = update(Throw(throwable)) @@ -706,9 +782,9 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ * value or an exception. setValue and setException are generally * more readable methods to use. * - * @throws ImmutableResult if the Promise is already populated + * @throws Promise.ImmutableResult if the Promise is already populated */ - def update(result: Try[A]): Unit = { + final def update(result: Try[A]): Unit = { updateIfEmpty(result) || { val current = Await.result(liftToTry) throw ImmutableResult(s"Result set multiple times. Value='$current', New='$result'") @@ -727,6 +803,13 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ */ @tailrec final def updateIfEmpty(result: Try[A]): Boolean = state match { + case s: Transforming[A] => + if (!cas(s, result)) updateIfEmpty(result) + else { + s.waitq.runInScheduler(result) + true + } + case waitq: WaitQueue[A] => if (!cas(waitq, result)) updateIfEmpty(result) else { @@ -734,47 +817,49 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ true } + case _: Try[A] /* Done */ => false + case s: Interruptible[A] => if (!cas(s, result)) updateIfEmpty(result) else { s.waitq.runInScheduler(result) true } - case s: Transforming[A] => - if (!cas(s, result)) updateIfEmpty(result) - else { - s.waitq.runInScheduler(result) - true - } + + case p: Promise[A] /* Linked */ => p.updateIfEmpty(result) + case s: Interrupted[A] => if (!cas(s, result)) updateIfEmpty(result) else { s.waitq.runInScheduler(result) true } - - case _: Try[A] /* Done */ => false - - case p: Promise[A] /* Linked */ => p.updateIfEmpty(result) } @tailrec - protected final def continue(k: K[A]): Unit = state match { - case waitq: WaitQueue[A] => - if (!cas(waitq, WaitQueue(k, waitq))) - continue(k) - case s: Interruptible[A] => - if (!cas(s, new Interruptible(WaitQueue(k, s.waitq), s.handler))) - continue(k) - case s: Transforming[A] => - if (!cas(s, new Transforming(WaitQueue(k, s.waitq), s.other))) - continue(k) - case s: Interrupted[A] => - if (!cas(s, new Interrupted(WaitQueue(k, s.waitq), s.signal))) - continue(k) - case v: Try[A] /* Done */ => k.runInScheduler(v) - case p: Promise[A] /* Linked */ => p.continue(k) - } + protected final def continue(wq: WaitQueue[A]): Unit = + if (wq ne WaitQueue.Empty) + state match { + case s: Transforming[A] => + if (!cas(s, new Transforming(WaitQueue.merge(wq, s.waitq), s.other))) + continue(wq) + + case waitq: WaitQueue[A] => + if (!cas(waitq, WaitQueue.merge(wq, waitq))) + continue(wq) + + case v: Try[A] /* Done */ => wq.runInScheduler(v) + + case s: Interruptible[A] => + if (!cas(s, new Interruptible(WaitQueue.merge(wq, s.waitq), s.handler, s.saved))) + continue(wq) + + case p: Promise[A] /* Linked */ => p.continue(wq) + + case s: Interrupted[A] => + if (!cas(s, new Interrupted(WaitQueue.merge(wq, s.waitq), s.signal))) + continue(wq) + } /** * Should only be called when this Promise has already been fulfilled @@ -785,7 +870,7 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ val target = p.compress() // due to the assumptions stated above regarding when this can be called, // there should never be a `cas` fail. - cas(p, target) + state = target target case _ => this @@ -796,39 +881,30 @@ class Promise[A] extends Future[A] with Promise.Responder[A] with Updatable[Try[ if (this eq target) return state match { - case waitq: WaitQueue[A] => - if (!cas(waitq, target)) link(target) - else target.continueAll(waitq) - - case s: Interruptible[A] => - if (!cas(s, target)) link(target) - else { - target.continueAll(s.waitq) - target.setInterruptHandler(s.handler) - } - case s: Transforming[A] => if (!cas(s, target)) link(target) - else { - target.continueAll(s.waitq) - target.forwardInterruptsTo(s.other) - } + else target.forwardInterruptsTo0(s.other, s.waitq) - case s: Interrupted[A] => - if (!cas(s, target)) link(target) - else { - target.continueAll(s.waitq) - target.raise(s.signal) - } + case waitq: WaitQueue[A] => + if (!cas(waitq, target)) link(target) + else target.continue(waitq) case value: Try[A] /* Done */ => if (!target.updateIfEmpty(value) && value != Await.result(target)) { throw new IllegalArgumentException("Cannot link two Done Promises with differing values") } + case s: Interruptible[A] => + if (!cas(s, target)) link(target) + else target.setInterruptHandler0(s.handler, s.waitq) + case p: Promise[A] /* Linked */ => if (cas(p, target)) p.link(target) else link(target) + + case s: Interrupted[A] => + if (!cas(s, target)) link(target) + else target.raise0(s.signal, s.waitq) } } diff --git a/util-core/src/main/scala/com/twitter/util/RandomSocket.scala b/util-core/src/main/scala/com/twitter/util/RandomSocket.scala index 33ec22ff8f..796a171955 100644 --- a/util-core/src/main/scala/com/twitter/util/RandomSocket.scala +++ b/util-core/src/main/scala/com/twitter/util/RandomSocket.scala @@ -5,7 +5,7 @@ import java.net.{InetSocketAddress, Socket} import scala.util.control.NonFatal /** - * A generator of random local [[java.net.InetSocketAddress]] objects with + * A generator of random local `java.net.InetSocketAddress` objects with * ephemeral ports. */ object RandomSocket { @@ -14,7 +14,7 @@ object RandomSocket { private[this] val ephemeralSocketAddress = localSocketOnPort(0) @deprecated("RandomSocket cannot ensure that the address is not in use.", "2014-11-13") - def apply() = nextAddress() + def apply(): InetSocketAddress = nextAddress() @deprecated("RandomSocket cannot ensure that the address is not in use.", "2014-11-13") def nextAddress(): InetSocketAddress = diff --git a/util-core/src/main/scala/com/twitter/util/ResourceTracker.scala b/util-core/src/main/scala/com/twitter/util/ResourceTracker.scala new file mode 100644 index 0000000000..c1199898f5 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/ResourceTracker.scala @@ -0,0 +1,115 @@ +package com.twitter.util + +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.atomic.LongAdder + +private[twitter] object ResourceTracker { + private[this] final val threadMxBean = java.lang.management.ManagementFactory.getThreadMXBean() + + /** + * An indicator if the running JDK supports querying thread statistics. + */ + final val threadCpuTimeSupported = threadMxBean.isCurrentThreadCpuTimeSupported + + /** + * Retreive the current thread's cpu time only if [[threadCpuTimeSupported]] is true. + * Otherwise, -1 is returned. Results may vary depending on the JDK's management implementation + */ + final def currentThreadCpuTime: Long = { + if (threadCpuTimeSupported) threadMxBean.getCurrentThreadCpuTime + else -1 + } + + /** + * Wrap the invocation of `f` with low-level resource usage that's pushed to the `tracker`. + */ + private[util] def wrapAndMeasureUsage[A, B]( + f: A => B, + tracker: ResourceTracker + ): A => B = { (result) => + { + tracker.incContinuations() + val cpuStart = threadMxBean.getCurrentThreadCpuTime + try f(result) + finally { + val cpuTime = threadMxBean.getCurrentThreadCpuTime - cpuStart + tracker.addCpuTime(cpuTime) + } + } + } + + /** + * Wrap the invocation of `f` with low-level resource usage that's pushed to the `tracker`. + */ + private[util] def wrapAndMeasureUsage[T]( + f: => T, + tracker: ResourceTracker + ): () => T = { () => + { + tracker.incContinuations() + val cpuStart = threadMxBean.getCurrentThreadCpuTime + try f + finally { + val cpuTime = threadMxBean.getCurrentThreadCpuTime - cpuStart + tracker.addCpuTime(cpuTime) + } + } + } + + // exposed for testing + private[util] def let[T](tracker: ResourceTracker)(f: => T): T = { + val oldCtx = Local.save() + val newCtx = oldCtx.setResourceTracker(tracker) + Local.restore(newCtx) + try f + finally Local.restore(oldCtx) + } + + // exposed for testing + private[util] def set(tracker: ResourceTracker): Unit = { + val oldCtx = Local.save() + val newCtx = oldCtx.setResourceTracker(tracker) + Local.restore(newCtx) + } + + // exposed for testing + private[util] def clear(): Unit = { + val oldCtx = Local.save() + val newCtx = oldCtx.removeResourceTracker() + if (newCtx ne oldCtx) Local.restore(newCtx) + } + + /** + * Track the execution of completed [[Future]]s within `f`. + */ + def apply[A](f: ResourceTracker => A) = { + val accumulator = new ResourceTracker() + let(accumulator)(f(accumulator)) + } +} + +/** + * A class used as the value for the local [[ResourceTracker#resourceTracker]] that's updated + * with low level resource usage as [[Promise]] continuations are executed. + * + * @note some measurements such as the cpu time relies on the JDK's implementation of the + * management interfaces. This means that the results may vary between JDK implementations. + */ +private[twitter] class ResourceTracker { + private[this] val cpuTime: LongAdder = new LongAdder() + private[this] val continuations: AtomicInteger = new AtomicInteger() + + private[util] def addCpuTime(time: Long): Unit = cpuTime.add(time) + private[util] def incContinuations(): Unit = continuations.getAndIncrement() + private[util] def addContinuations(count: Int): Unit = continuations.getAndAdd(count) + + /** + * Retrieve the total accumulated CPU time + */ + def totalCpuTime(): Long = cpuTime.sum() + + /** + * Retrieve the total number of executed continuations + */ + def numContinuations(): Int = continuations.get() +} diff --git a/util-core/src/main/scala/com/twitter/util/SchedulerWorkQueueFiber.scala b/util-core/src/main/scala/com/twitter/util/SchedulerWorkQueueFiber.scala new file mode 100644 index 0000000000..a0f3e8e6bf --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/SchedulerWorkQueueFiber.scala @@ -0,0 +1,31 @@ +package com.twitter.util + +import com.twitter.concurrent.Scheduler + +/** + * A [[WorkQueueFiber]] implementation that submits its work queue to a scheduler. + * + * This allows us to gain the benefits of the WorkQueueFiber, namely the bundling of work + * into Runnables of multiple tasks, while still utilizing the global scheduler. This + * implementation of the WorkQueueFiber is meant for servers that are not well suited + * for offloading tasks into a thread pool for execution. + * + * @param scheduler the task scheduler that will execute tasks. + * @param fiberMetrics telemetry used for tracking the execution of this fiber, and likely + * other fibers as well. + */ +private[twitter] final class SchedulerWorkQueueFiber( + scheduler: Scheduler, + fiberMetrics: WorkQueueFiber.FiberMetrics) + extends WorkQueueFiber(fiberMetrics) { + + /** + * Submit the work to the underlying scheduler for execution. + */ + override protected def schedulerSubmit(r: Runnable): Unit = scheduler.submit(r) + + /** + * Flush the underlying scheduler. + */ + override protected def schedulerFlush(): Unit = scheduler.flush() +} diff --git a/util-core/src/main/scala/com/twitter/util/Signal.scala b/util-core/src/main/scala/com/twitter/util/Signal.scala index cfe25569f0..2a945361c1 100644 --- a/util-core/src/main/scala/com/twitter/util/Signal.scala +++ b/util-core/src/main/scala/com/twitter/util/Signal.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -21,9 +21,9 @@ import java.lang.reflect.{InvocationHandler, Method, Proxy} import scala.collection.{Map, Set, mutable} object SignalHandlerFactory { - def apply() = { + def apply(): Option[SunSignalHandler] = { // only one actual implementation for now - SunSignalHandler.instantiate + SunSignalHandler.instantiate() } } @@ -32,7 +32,7 @@ trait SignalHandler { } object SunSignalHandler { - def instantiate() = { + def instantiate(): Option[SunSignalHandler] = { try { Class.forName("sun.misc.Signal") Some(new SunSignalHandler()) @@ -58,9 +58,7 @@ class SunSignalHandler extends SignalHandler { new InvocationHandler { def invoke(proxy: Object, method: Method, args: Array[Object]) = { if (method.getName() == "handle") { - handlers(signal).foreach { x => - x(nameMethod.invoke(args(0)).asInstanceOf[String]) - } + handlers(signal).foreach { x => x(nameMethod.invoke(args(0)).asInstanceOf[String]) } } null } diff --git a/util-core/src/main/scala/com/twitter/util/SlowProbeProxyTimer.scala b/util-core/src/main/scala/com/twitter/util/SlowProbeProxyTimer.scala index 45d5302175..1dbe4ef280 100644 --- a/util-core/src/main/scala/com/twitter/util/SlowProbeProxyTimer.scala +++ b/util-core/src/main/scala/com/twitter/util/SlowProbeProxyTimer.scala @@ -48,7 +48,9 @@ abstract class SlowProbeProxyTimer(maxRuntime: Duration) extends ProxyTimer { override protected def schedulePeriodically( when: Time, period: Duration - )(f: => Unit): TimerTask = { + )( + f: => Unit + ): TimerTask = { checkSlowTask() self.schedule(when, period)(meterTask(f)) } diff --git a/util-core/src/main/scala/com/twitter/util/StateMachine.scala b/util-core/src/main/scala/com/twitter/util/StateMachine.scala index f138baa9e1..434f87fa95 100644 --- a/util-core/src/main/scala/com/twitter/util/StateMachine.scala +++ b/util-core/src/main/scala/com/twitter/util/StateMachine.scala @@ -11,7 +11,7 @@ trait StateMachine { protected abstract class State protected var state: State = _ - protected def transition[A](command: String)(f: PartialFunction[State, A]) = synchronized { + protected def transition[A](command: String)(f: PartialFunction[State, A]): A = synchronized { if (f.isDefinedAt(state)) { f(state) } else { diff --git a/util-core/src/main/scala/com/twitter/util/Stopwatch.scala b/util-core/src/main/scala/com/twitter/util/Stopwatch.scala index 1b39a4b8db..9f733e3506 100644 --- a/util-core/src/main/scala/com/twitter/util/Stopwatch.scala +++ b/util-core/src/main/scala/com/twitter/util/Stopwatch.scala @@ -39,11 +39,15 @@ object Stopwatch extends Stopwatch { * A function that returns a Long that can be used for measuring elapsed time * in nanoseconds, that is `Time` manipulation compatible. * - * Useful for testing, but should not be used in production, since it uses a - * non-monotonic time under the hood, and entails a few allocations. For + * Useful for testing, but should not be used in a tight loop in production, + * since it entails a few allocations and does a [[Local]] lookup. For * production, see `systemNanos`. */ - val timeNanos: () => Long = () => Time.now.inNanoseconds + val timeNanos: () => Long = () => + Time.localGetTime() match { + case Some(local) => local().inNanoseconds + case None => systemNanos() + } /** * A function that returns a Long that can be used for measuring elapsed time @@ -59,11 +63,15 @@ object Stopwatch extends Stopwatch { * A function that returns a Long that can be used for measuring elapsed time * in microseconds, that is `Time` manipulation compatible. * - * Useful for testing, but should not be used in production, since it uses a - * non-monotonic time under the hood, and entails a few allocations. For + * Useful for testing, but should not be used in a tight loop in production, + * since it entails a few allocations and does a [[Local]] lookup. For * production, see `systemMicros`. */ - val timeMicros: () => Long = () => Time.now.inMicroseconds + val timeMicros: () => Long = () => + Time.localGetTime() match { + case Some(local) => local().inMicroseconds + case None => systemMicros() + } /** * A function that returns a Long that can be used for measuring elapsed time @@ -79,21 +87,23 @@ object Stopwatch extends Stopwatch { * A function that returns a Long that can be used for measuring elapsed time * in milliseconds, that is `Time` manipulation compatible. * - * Useful for testing, but should not be used in production, since it uses a - * non-monotonic time under the hood, and entails a few allocations. For + * Useful for testing, but should not be used in a tight loop in production, + * since it entails a few allocations and does a [[Local]] lookup. For * production, see `systemMillis`. */ - val timeMillis: () => Long = () => Time.now.inMilliseconds + val timeMillis: () => Long = () => + Time.localGetTime() match { + case Some(local) => local().inMilliseconds + case None => systemMillis() + } def start(): Elapsed = Time.localGetTime() match { case Some(local) => val startAt: Time = local() - () => - local() - startAt + () => local() - startAt case None => val startAt: Long = systemNanos() - () => - Duration.fromNanoseconds(systemNanos() - startAt) + () => Duration.fromNanoseconds(systemNanos() - startAt) } /** @@ -112,7 +122,7 @@ object Stopwatch extends Stopwatch { object Stopwatches { import Stopwatch.Elapsed - /** @see [[Stopwatch.start()]] */ + /** @see [[Stopwatch.start]] */ def start(): Elapsed = Stopwatch.start() @@ -149,7 +159,7 @@ object Stopwatches { * A trivial implementation of [[Stopwatch]] for use as a null * object. * - * All calls to [[Stopwatch.start()]] return an [[Stopwatch.Elapsed]] + * All calls to [[Stopwatch.start]] return an [[Stopwatch.Elapsed]] * instance that always returns [[Duration.Bottom]] for its elapsed time. * * @see `NilStopwatch.get` for accessing this instance from Java. diff --git a/util-core/src/main/scala/com/twitter/util/StorageUnit.scala b/util-core/src/main/scala/com/twitter/util/StorageUnit.scala index 01620d4111..f1d8cf2a8d 100644 --- a/util-core/src/main/scala/com/twitter/util/StorageUnit.scala +++ b/util-core/src/main/scala/com/twitter/util/StorageUnit.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -31,8 +31,8 @@ object StorageUnit { def fromExabytes(exabytes: Long): StorageUnit = new StorageUnit(exabytes * 1024 * 1024 * 1024 * 1024 * 1024 * 1024) - val infinite = new StorageUnit(Long.MaxValue) - val zero = new StorageUnit(0) + val infinite: StorageUnit = new StorageUnit(Long.MaxValue) + val zero: StorageUnit = new StorageUnit(0) private def factor(s: String): Long = { var lower = s.toLowerCase @@ -51,6 +51,13 @@ object StorageUnit { } } + // We define this internally since `signum` is deprecated since Scala 2.13 in favor of `sig`. + // Lets define this simple method here until support for Scala 2.12 is dropped. + private def sig(bytes: Long): Int = + if (bytes > 0) 1 + else if (bytes < 0) -1 + else 0 + /** * Note, this can cause overflows of the Long used to represent the * number of bytes. @@ -69,11 +76,11 @@ object StorageUnit { /** * Representation of storage units. * - * Use either `StorageUnit.fromX` or [[com.twitter.conversions.storage implicit conversions]] + * Use either `StorageUnit.fromX` or `com.twitter.conversions.storage` `implicit conversions` * from `Long` and `Int` to construct instances. * * {{{ - * import com.twitter.conversions.storage._ + * import com.twitter.conversions.StorageUnitOps._ * import com.twitter.util.StorageUnit * * val size: StorageUnit = 10.kilobytes @@ -87,6 +94,8 @@ object StorageUnit { * number of bytes. */ class StorageUnit(val bytes: Long) extends Ordered[StorageUnit] { + import StorageUnit.sig + def inBytes: Long = bytes def inKilobytes: Long = bytes / 1024L def inMegabytes: Long = bytes / (1024L * 1024) @@ -126,7 +135,7 @@ class StorageUnit(val bytes: Long) extends Ordered[StorageUnit] { def max(other: StorageUnit): StorageUnit = if (this > other) this else other - override def toString(): String = inBytes + ".bytes" + override def toString(): String = s"$inBytes.bytes" def toHuman(): String = { val prefix = "KMGTPE" @@ -139,7 +148,8 @@ class StorageUnit(val bytes: Long) extends Ordered[StorageUnit] { if (prefixIndex < 0) { "%d B".format(bytes) } else { - "%.1f %ciB".formatLocal(Locale.ENGLISH, display * bytes.signum, prefix.charAt(prefixIndex)) + "%.1f %ciB".formatLocal(Locale.ENGLISH, display * sig(bytes), prefix.charAt(prefixIndex)) } } + } diff --git a/util-core/src/main/scala/com/twitter/util/Throwables.scala b/util-core/src/main/scala/com/twitter/util/Throwables.scala index 67bdf8f010..9057a172a5 100644 --- a/util-core/src/main/scala/com/twitter/util/Throwables.scala +++ b/util-core/src/main/scala/com/twitter/util/Throwables.scala @@ -1,7 +1,7 @@ package com.twitter.util import scala.annotation.tailrec -import scala.collection.mutable.ArrayBuffer +import scala.collection.mutable object Throwables { @@ -27,17 +27,27 @@ object Throwables { * classname strings. */ def mkString(ex: Throwable): Seq[String] = { - @tailrec def rec(ex: Throwable, buf: ArrayBuffer[String]): Seq[String] = { + @tailrec def rec(ex: Throwable, buf: List[String]): Seq[String] = { if (ex eq null) - buf.result + buf.reverse else - rec(ex.getCause, buf += ex.getClass.getName) + rec(ex.getCause, ex.getClass.getName :: buf) } - rec(ex, ArrayBuffer.empty) + rec(ex, Nil) } object RootCause { def unapply(e: Throwable): Option[Throwable] = Option(e.getCause) + + def nested(ex: Throwable): Throwable = { + var rootException = ex + val exceptions = mutable.HashSet[Throwable]() + while (rootException.getCause != null && !exceptions.contains(rootException)) { + exceptions.add(rootException) + rootException = rootException.getCause + } + rootException + } } } diff --git a/util-core/src/main/scala/com/twitter/util/Time.scala b/util-core/src/main/scala/com/twitter/util/Time.scala index b94e583d4d..cedde8064b 100644 --- a/util-core/src/main/scala/com/twitter/util/Time.scala +++ b/util-core/src/main/scala/com/twitter/util/Time.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -17,8 +17,17 @@ package com.twitter.util import java.io.Serializable +import java.text.ParseException +import java.time.ZonedDateTime +import java.time.Clock +import java.time.Instant +import java.time.OffsetDateTime +import java.time.ZoneOffset +import java.time.format.DateTimeParseException import java.util.concurrent.TimeUnit -import java.util.{Date, Locale, TimeZone} +import java.util.Date +import java.util.Locale +import java.util.TimeZone /** * @define now @@ -96,6 +105,12 @@ trait TimeLikeOps[This <: TimeLike[This]] { def fromSeconds(seconds: Int): This = fromMilliseconds(1000L * seconds) + def fromMinutes(minutes: Int): This = fromMilliseconds(60L * 1000L * minutes) + + def fromHours(hours: Int): This = fromMilliseconds(60L * 60L * 1000L * hours) + + def fromDays(days: Int): This = fromMilliseconds(24L * 3600L * 1000L * days) + /** * Make a new `This` from the given number of seconds. * Because this method takes a Double, it can represent values less than a second. @@ -136,8 +151,7 @@ trait TimeLikeOps[This <: TimeLike[This]] { * Overflows are also handled like doubles. */ trait TimeLike[This <: TimeLike[This]] extends Ordered[This] { self: This => - protected val ops: TimeLikeOps[This] - import ops._ + protected def ops: TimeLikeOps[This] /** The `TimeLike`'s value in nanoseconds. */ def inNanoseconds: Long @@ -156,20 +170,17 @@ trait TimeLike[This <: TimeLike[This]] extends Ordered[This] { self: This => def inMillis: Long = inMilliseconds // (Backwards compatibility) /** - * Returns a value/`TimeUnit` pair; attempting to return coarser - * grained values if possible (specifically: `TimeUnit.SECONDS` or - * `TimeUnit.MILLISECONDS`) before resorting to the default + * Returns a value/`TimeUnit` pair; attempting to return coarser grained + * values if possible (specifically: `TimeUnit.MINUTES`, `TimeUnit.SECONDS` + * or `TimeUnit.MILLISECONDS`) before resorting to the default * `TimeUnit.NANOSECONDS`. */ - def inTimeUnit: (Long, TimeUnit) = { + def inTimeUnit: (Long, TimeUnit) = inNanoseconds match { // allow for APIs that may treat TimeUnit differently if measured in very tiny units. - if (inNanoseconds % Duration.NanosPerSecond == 0) { - (inSeconds, TimeUnit.SECONDS) - } else if (inNanoseconds % Duration.NanosPerMillisecond == 0) { - (inMilliseconds, TimeUnit.MILLISECONDS) - } else { - (inNanoseconds, TimeUnit.NANOSECONDS) - } + case ns if ns % Duration.NanosPerMinute == 0 => (inMinutes, TimeUnit.MINUTES) + case ns if ns % Duration.NanosPerSecond == 0 => (inSeconds, TimeUnit.SECONDS) + case ns if ns % Duration.NanosPerMillisecond == 0 => (inMilliseconds, TimeUnit.MILLISECONDS) + case ns => (ns, TimeUnit.NANOSECONDS) } /** @@ -182,14 +193,14 @@ trait TimeLike[This <: TimeLike[This]] extends Ordered[This] { self: This => def addNanos(a: Long, b: Long): This = { val c = a + b if (((a ^ c) & (b ^ c)) < 0) - if (b < 0) Bottom else Top - else fromNanoseconds(c) + if (b < 0) ops.Bottom else ops.Top + else ops.fromNanoseconds(c) } delta match { - case Duration.Top => Top - case Duration.Bottom => Bottom - case Duration.Undefined => Undefined + case Duration.Top => ops.Top + case Duration.Bottom => ops.Bottom + case Duration.Undefined => ops.Undefined case ns => addNanos(inNanoseconds, ns.inNanoseconds) } } @@ -223,11 +234,13 @@ trait TimeLike[This <: TimeLike[This]] extends Ordered[This] { self: This => * results because of timezones. */ def floor(increment: Duration): This = (this, increment) match { - case (num, ns) if num.isZero && ns.isZero => Undefined - case (num, ns) if num.isFinite && ns.isZero => if (num.inNanoseconds < 0) Bottom else Top - case (num, denom) if num.isFinite && denom.isFinite => fromNanoseconds((num.inNanoseconds / denom.inNanoseconds) * denom.inNanoseconds) + case (num, ns) if num.isZero && ns.isZero => ops.Undefined + case (num, ns) if num.isFinite && ns.isZero => + if (num.inNanoseconds < 0) ops.Bottom else ops.Top + case (num, denom) if num.isFinite && denom.isFinite => + ops.fromNanoseconds((num.inNanoseconds / denom.inNanoseconds) * denom.inNanoseconds) case (self, n) if n.isFinite => self - case (_, _) => Undefined + case (_, _) => ops.Undefined } def max(that: This): This = @@ -237,15 +250,15 @@ trait TimeLike[This <: TimeLike[This]] extends Ordered[This] { self: This => if ((this compare that) < 0) this else that def compare(that: This): Int = - if ((that eq Top) || (that eq Undefined)) -1 - else if (that eq Bottom) 1 + if ((that eq ops.Top) || (that eq ops.Undefined)) -1 + else if (that eq ops.Bottom) 1 else if (inNanoseconds < that.inNanoseconds) -1 else if (inNanoseconds > that.inNanoseconds) 1 else 0 /** Equality within `maxDelta` */ def moreOrLessEquals(other: This, maxDelta: Duration): Boolean = - (other ne Undefined) && ((this == other) || (this diff other).abs <= maxDelta) + (other ne ops.Undefined) && ((this == other) || (this diff other).abs <= maxDelta) } /** @@ -283,17 +296,23 @@ private[util] object TimeBox { * * @see [[Duration]] for intervals between two points in time. * @see [[Stopwatch]] for measuring elapsed time. - * @see [[TimeFormat]] for converting to and from `String` representations. + * @see [[TimeFormatter]] for converting to and from `String` representations. */ object Time extends TimeLikeOps[Time] { def fromNanoseconds(nanoseconds: Long): Time = new Time(nanoseconds) + def fromInstant(instant: Instant): Time = new Time( + instant.getEpochSecond * Duration.NanosPerSecond + instant.getNano) + // This is needed for Java compatibility. override def fromFractionalSeconds(seconds: Double): Time = super.fromFractionalSeconds(seconds) override def fromSeconds(seconds: Int): Time = super.fromSeconds(seconds) + override def fromMinutes(minutes: Int): Time = super.fromMinutes(minutes) override def fromMilliseconds(millis: Long): Time = super.fromMilliseconds(millis) override def fromMicroseconds(micros: Long): Time = super.fromMicroseconds(micros) + private[this] val SystemClock = Clock.systemUTC() + /** * Time `Top` is greater than any other definable time, and is * equal only to itself. It may be used as a sentinel value, @@ -319,7 +338,7 @@ object Time extends TimeLikeOps[Time] { override def diff(that: Time): Duration = that match { case Top | Undefined => Duration.Undefined - case other => Duration.Top + case _ => Duration.Top } override def isFinite: Boolean = false @@ -349,7 +368,7 @@ object Time extends TimeLikeOps[Time] { override def diff(that: Time): Duration = that match { case Bottom | Undefined => Duration.Undefined - case other => Duration.Bottom + case _ => Duration.Bottom } override def isFinite: Boolean = false @@ -392,30 +411,109 @@ object Time extends TimeLikeOps[Time] { } } + /** + * Returns the current [[Time]] with, at best, nanosecond precision. + * + * Note that nanosecond precision is only available in JDK versions + * greater than JDK8. In JDK8 this API has the same precision as + * [[Time#now]] and `System#currentTimeMillis`. In JDK9+ this will + * change and all timestamps taken from this API will have nanosecond + * resolution. + * + * Note that returned values are '''not''' monotonic. This + * means it is not suitable for measuring durations where + * the clock cannot drift backwards. If monotonicity is desired + * prefer a monotonic [[Stopwatch]]. + * + */ + def nowNanoPrecision: Time = { + localGetTime() match { + case None => + val rightNow = Instant.now(SystemClock) + fromNanoseconds((rightNow.getEpochSecond * Duration.NanosPerSecond) + rightNow.getNano) + case Some(f) => + f() + } + } + /** * The unix epoch. Times are measured relative to this. */ val epoch: Time = fromNanoseconds(0L) - private val defaultFormat = new TimeFormat("yyyy-MM-dd HH:mm:ss Z") - private val rssFormat = new TimeFormat("E, dd MMM yyyy HH:mm:ss Z") + private val defaultFormat = TimeFormatter("yyyy-MM-dd HH:mm:ss Z") + + // The default new parsing pattern is different from the default formatting pattern. There + // are two reasons for this, both of them designed to make parsing backwards-compatible with + // TimeFormat's SimpleDateFormat-based solution. One is that DateTimeFormatter pattern Z doesn't + // recognize textual zone IDs and z recognizes mostly just textual zone IDs, so we add both of + // them as optional. (Unfortunately, we can't enforce exactly one to be accepted, so it is + // possible to parse a string ending with a space and no time zone, or two time zones in which + // case last zone wins.) See https://bugs.openjdk.org/browse/JDK-8294710 + // Another reason is that old SimpleDateFormat accepts single-digits for months, days, hours, + // minutes, and seconds with two-character patterns, while DateTimeFormatter doesn't. So for + // formatting we output canonical two-digit forms, but for parsing we accept single-digit forms. + private[this] val defaultParser = + versionedParser(oldPattern = "yyyy-MM-dd HH:mm:ss Z", newPattern = "yyyy-M-d H:m:s [Z][z]") + + // The new RSS datetime pattern leverages the optional elements feature to correctly specify that + // day-of-week and second-of-minute are optional in RFC-1123. It also uses the same [Z][z] pattern + // as default parse. + private[this] val rssParser = versionedParser( + oldPattern = "E, d MMM yyyy HH:mm:ss Z", + newPattern = "[E, ]d MMM yyyy HH:mm[:ss] [Z][z]") + + // This creates a parser function that either uses the deprecated TimeFormat or the new + // TimeFormatter, depending on whether the system executes on a 1.8 Java runtime or not. This is + // We still use the old SimpleDateFormat-based format for parsing Time.at/fromRss when running on + // Java 1.8 as a workaround for a bug with correct timezone parsing with the new java.time API + // that is fixed in Java 11 and later, see https://bugs.openjdk.org/browse/JDK-8294853. + private def versionedParser(oldPattern: String, newPattern: String): String => Time = { + if (System.getProperty("java.version").startsWith("1.8")) { + val tf = new TimeFormat(oldPattern) + tf.parse + } else { + // Setting Locale.US will provide the expected textual values for day-of-week and month-of-year. + val tf = TimeFormatter(newPattern, Locale.US) + s => + try { + tf.parse(s) + } catch { + case e: DateTimeParseException => throw translateException(e) + } + } + } /** * Note, this should only ever be updated by methods used for testing. */ - private[util] val localGetTime = new Local[() => Time] - private[util] val localGetTimer = new Local[MockTimer] + private[util] val localGetTime: Local[() => Time] = new Local[() => Time] + private[util] val localGetTimer: Local[MockTimer] = new Local[MockTimer] /** - * Creates a [[Time]] instance of the given [[Date]]. + * Creates a [[Time]] instance of the given `Date`. */ def apply(date: Date): Time = fromMilliseconds(date.getTime) + private[this] def translateException(e: DateTimeParseException): ParseException = { + // Methods calling this used to throw java.text.ParseException before, + // and some callers might rely on that. + val pe = new ParseException(e.getMessage, e.getErrorIndex) + pe.initCause(e) + pe + } + /** * Creates a [[Time]] instance at the given `datetime` string in the * "yyyy-MM-dd HH:mm:ss Z" format. + * + * {{{ + * val time = Time.at("2019-05-12 00:00:00 -0800") + * time.toString // 2019-05-12 08:00:00 +0000 + * }}} */ - def at(datetime: String): Time = defaultFormat.parse(datetime) + def at(datetime: String): Time = + defaultParser(datetime) /** * Execute body with the time function replaced by `timeFunction` @@ -491,7 +589,8 @@ object Time extends TimeLikeOps[Time] { /** * Returns the Time parsed from a string in RSS format. Eg: "Wed, 15 Jun 2005 19:00:00 GMT" */ - def fromRss(rss: String): Time = rssFormat.parse(rss) + def fromRss(rss: String): Time = + rssParser(rss) } trait TimeControl { @@ -502,13 +601,29 @@ trait TimeControl { /** * A thread-safe wrapper around a SimpleDateFormat object. * + * This class is deprecated in favor of [[TimeFormatter]]. When moving from this class to + * [[TimeFormatter]], be mindful that there are some subtle differences in patterns between + * [[java.text.SimpleDateFormat]] used by this class and [[DateTimeFormatter]] used by [[TimeFormatter]]. + * Pattern character 'u' has a different meaning. Character 'a' cannot be repeated, and numerals for + * day of week produced by 'e' and 'c' can count from either Sunday or Monday depending on locale. + * Time zones might be parsed differently, and [[DateTimeFormatter]] defines several different time + * zone pattern characters with various behaviors. Consult the Javadoc for [[java.text.SimpleDateFormat]] + * and [[DateTimeFormatter]] for details. Finally, the parse method in this class throws a + * [[java.text.ParseException]] while the one in [[TimeFormatter]] throws [[DateTimeParseException]]. + * If you rely on catching exceptions from parsing, you will need to change the type of the exception. + * Note that built in [[Time.at]] and [[Time.fromRss]] keep throwing [[java.text.ParseException]] for + * backwards compatibility. + * * The default timezone is UTC. */ +@deprecated( + message = + "Use TimeFormatter instead; it uses modern time APIs and requires no synchronization. " + + "See TimeFormat documentation for information on pattern differences.") class TimeFormat( pattern: String, locale: Option[Locale], - timezone: TimeZone = TimeZone.getTimeZone("UTC") -) { + timezone: TimeZone = TimeZone.getTimeZone("UTC")) { // jdk6 and jdk7 pick up the default locale differently in SimpleDateFormat, // so we can't rely on Locale.getDefault here. @@ -550,12 +665,13 @@ class TimeFormat( * * @see [[Duration]] for intervals between two points in time. * @see [[Stopwatch]] for measuring elapsed time. - * @see [[TimeFormat]] for converting to and from `String` representations. + * @see [[TimeFormatter]] for converting to and from `String` representations. */ -sealed class Time private[util] (protected val nanos: Long) extends { - protected val ops = Time -} with TimeLike[Time] with Serializable { - import ops._ +sealed class Time private[util] (protected val nanos: Long) + extends TimeLike[Time] + with Serializable { + protected def ops: Time.type = Time + import Time._ def inNanoseconds: Long = nanos @@ -583,11 +699,17 @@ sealed class Time private[util] (protected val nanos: Long) extends { /** * Formats this Time according to the given SimpleDateFormat pattern. */ + @deprecated( + message = + "Use TimeFormatter instead. See TimeFormat documentation for information on pattern differences.") def format(pattern: String): String = new TimeFormat(pattern).format(this) /** * Formats this Time according to the given SimpleDateFormat pattern and locale. */ + @deprecated( + message = + "Use TimeFormatter instead. See TimeFormat documentation for information on pattern differences.") def format(pattern: String, locale: Locale): String = new TimeFormat(pattern, Some(locale)).format(this) @@ -650,6 +772,22 @@ sealed class Time private[util] (protected val nanos: Long) extends { */ def toDate: Date = new Date(inMillis) + /** + * Converts this Time object to a java.time.Instant + */ + def toInstant: Instant = + Instant.ofEpochSecond(0L, nanos) + + /** + * Converts this Time object to a java.time.ZonedDateTime with ZoneId of UTC + */ + def toZonedDateTime: ZonedDateTime = ZonedDateTime.ofInstant(toInstant, ZoneOffset.UTC) + + /** + * Converts this Time object to a java.time.OffsetDateTime with ZoneId of UTC + */ + def toOffsetDateTime: OffsetDateTime = OffsetDateTime.ofInstant(toInstant, ZoneOffset.UTC) + private def writeReplace(): Object = TimeBox.Finite(inNanoseconds) /** diff --git a/util-core/src/main/scala/com/twitter/util/TimeFormatter.scala b/util-core/src/main/scala/com/twitter/util/TimeFormatter.scala new file mode 100644 index 0000000000..3db7dc5add --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/TimeFormatter.scala @@ -0,0 +1,81 @@ +/* + * Copyright 2022 Twitter Inc. + * + * 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 + * + * https://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 com.twitter.util + +import java.time.Instant +import java.time.ZoneId +import java.time.ZoneOffset +import java.time.format.DateTimeFormatter +import java.time.format.DateTimeFormatterBuilder +import java.time.temporal.ChronoField +import java.util.Locale + +object TimeFormatter { + + /** + * Creates a new [[TimeFormatter]] based on a pattern and optionally a locale and a time zone. + * @param pattern the [[DateTimeFormatter]] compatible pattern + * @param locale locale, defaults to platform default locale for formatting + * @param zoneId time zone, defaults to UTC + * @return a new [[TimeFormatter]] for the specified pattern, locale, and time zone. + */ + def apply( + pattern: String, + locale: Locale = Locale.getDefault(Locale.Category.FORMAT), + zoneId: ZoneId = ZoneOffset.UTC + ): TimeFormatter = { + val (has_y, has_M, has_d, has_H) = TwitterDateFormat.validatePattern(pattern, newFormat = true) + val builder = new DateTimeFormatterBuilder() + .parseLenient() + .parseCaseInsensitive() + .appendPattern(pattern) + // These defaults match those that SimpleDateFormat would use during + // parsing when these fields are not specified. [[Instant.from]] can + // produce an Instant from parse results if at least all the fields up to + // HOUR_OF_DAY are present. + if (!has_y) builder.parseDefaulting(ChronoField.YEAR_OF_ERA, 1970) + if (!has_M) builder.parseDefaulting(ChronoField.MONTH_OF_YEAR, 1) + if (!has_d) builder.parseDefaulting(ChronoField.DAY_OF_MONTH, 1) + if (!has_H) builder.parseDefaulting(ChronoField.HOUR_OF_DAY, 0) + new TimeFormatter(builder.toFormatter(locale).withZone(zoneId)) + } +} + +/** + * A parser and formatter for Time instances. Instances of this class are immutable and threadsafe. + * Use [[TimeFormatter$.apply]] to construct instances from [[DateTimeFormatter]]-compatible pattern + * strings. + */ +class TimeFormatter private[util] (format: DateTimeFormatter) { + + /** + * Parses a string into a [[Time]]. + * @param str the string to parse + * @return the parsed time object + * @throws java.time.format.DateTimeParseException if the string can not be parsed + */ + def parse(str: String): Time = + Time.fromInstant(Instant.from(format.parse(str))) + + /** + * Formats a [[Time]] object into a string. + * @param time the time object to format + * @return a string with formatted representation of the time object + */ + def format(time: Time): String = + format.format(time.toInstant) +} diff --git a/util-core/src/main/scala/com/twitter/util/Timer.scala b/util-core/src/main/scala/com/twitter/util/Timer.scala index 325e619018..6c2f0ecc57 100644 --- a/util-core/src/main/scala/com/twitter/util/Timer.scala +++ b/util-core/src/main/scala/com/twitter/util/Timer.scala @@ -1,7 +1,7 @@ package com.twitter.util import com.twitter.concurrent.NamedPoolThreadFactory -import java.util.concurrent._ +import java.util.concurrent.{Future => _, _} import scala.collection.mutable.ArrayBuffer import scala.util.control @@ -188,23 +188,23 @@ class ReferenceCountingTimer(factory: () => Timer) extends ProxyTimer with Refer } /** - * A [[Timer]] that is implemented via a [[java.util.Timer]]. + * A [[Timer]] that is implemented via a `java.util.Timer`. * * This timer has millisecond granularity. * * If your code has a reasonably high throughput of task scheduling * and can trade off some precision of when tasks run, * [[https://twitter.github.io/finagle/ Finagle]] has a higher throughput - * [[https://github.com/twitter/finagle/blob/master/finagle-core/src/main/scala/com/twitter/finagle/util/DefaultTimer.scala + * [[https://github.com/twitter/finagle/blob/release/finagle-core/src/main/scala/com/twitter/finagle/util/DefaultTimer.scala * hashed-wheel implementation]]. * * @note Due to the implementation using a single `Thread`, be wary of * scheduling tasks that take a long time to execute (e.g. blocking IO) * as this blocks other tasks from running at their scheduled time. * - * @param isDaemon whether or not the associated [[Thread]] should run + * @param isDaemon whether or not the associated `Thread` should run * as a daemon. - * @param name used as the name of the associated [[Thread]] when specified. + * @param name used as the name of the associated `Thread` when specified. */ class JavaTimer(isDaemon: Boolean, name: Option[String]) extends Timer { def this(isDaemon: Boolean) = this(isDaemon, None) @@ -272,8 +272,8 @@ class JavaTimer(isDaemon: Boolean, name: Option[String]) extends Timer { class ScheduledThreadPoolTimer( poolSize: Int, threadFactory: ThreadFactory, - rejectedExecutionHandler: Option[RejectedExecutionHandler] -) extends Timer { + rejectedExecutionHandler: Option[RejectedExecutionHandler]) + extends Timer { def this(poolSize: Int, threadFactory: ThreadFactory) = this(poolSize, threadFactory, None) @@ -347,7 +347,7 @@ class MockTimer extends Timer { // These are weird semantics admittedly, but there may // be a bunch of tests that rely on them already. case class Task(var when: Time, runner: () => Unit) extends TimerTask { - var isCancelled = false + var isCancelled: Boolean = false def cancel(): Unit = MockTimer.this.synchronized { isCancelled = true @@ -357,18 +357,16 @@ class MockTimer extends Timer { } } - var isStopped = false - var tasks = ArrayBuffer[Task]() - var nCancelled = 0 + var isStopped: Boolean = false + var tasks: ArrayBuffer[Task] = ArrayBuffer[Task]() + var nCancelled: Int = 0 def tick(): Unit = synchronized { if (isStopped) throw new IllegalStateException("timer is stopped already") val now = Time.now - val (toRun, toQueue) = tasks.partition { task => - task.when <= now - } + val (toRun, toQueue) = tasks.partition { task => task.when <= now } tasks = toQueue toRun.filterNot(_.isCancelled).foreach(_.runner()) } diff --git a/util-core/src/main/scala/com/twitter/util/Try.scala b/util-core/src/main/scala/com/twitter/util/Try.scala index b9c898d350..252ca36b24 100644 --- a/util-core/src/main/scala/com/twitter/util/Try.scala +++ b/util-core/src/main/scala/com/twitter/util/Try.scala @@ -12,8 +12,19 @@ import scala.util.{Success, Failure} object Try { case class PredicateDoesNotObtain() extends Exception() + /** + * A constant `Try` that returns `Unit`. + */ + val Unit: Try[Unit] = Try(()) + + /** + * A constant `Try` that returns `Void`. + */ + val Void: Try[Void] = Try(null: Void) + def apply[R](r: => R): Try[R] = { - try { Return(r) } catch { + try { Return(r) } + catch { case NonFatal(e) => Throw(e) } } @@ -47,9 +58,7 @@ object Try { if (ts.isEmpty) Return(Seq.empty[A]) else Try { - ts map { t => - t() - } + ts map { t => t() } } } @@ -118,7 +127,7 @@ sealed abstract class Try[+R] { /** * Returns the value from this Return or the given argument if this is a Throw. */ - def getOrElse[R2 >: R](default: => R2) = if (isReturn) apply() else default + def getOrElse[R2 >: R](default: => R2): R2 = if (isReturn) apply() else default /** * Returns the value from this Return or throws the exception if this is a Throw @@ -129,7 +138,7 @@ sealed abstract class Try[+R] { * Returns the value from this Return or throws the exception if this is a Throw. * Alias for apply() */ - def get() = apply() + def get(): R = apply() /** * Applies the given function f if this is a Result. @@ -197,14 +206,12 @@ sealed abstract class Try[+R] { * chained `this` as in `respond`. */ def ensure(f: => Unit): Try[R] = - respond { _ => - f - } + respond { _ => f } /** * Returns None if this is a Throw or a Some containing the value if this is a Return */ - def toOption = if (isReturn) Some(apply()) else None + def toOption: Option[R] = if (isReturn) Some(apply()) else None /** * Invokes the given closure when the value is available. Returns @@ -226,7 +233,7 @@ sealed abstract class Try[+R] { * Returns the given function applied to the value from this Return or returns this if this is a Throw. * Alias for flatMap */ - def andThen[R2](f: R => Try[R2]) = flatMap(f) + def andThen[R2](f: R => Try[R2]): Try[R2] = flatMap(f) def flatten[T](implicit ev: R <:< Try[T]): Try[T] @@ -239,10 +246,10 @@ object Throw { final case class Throw[+R](e: Throwable) extends Try[R] { def asScala: scala.util.Try[R] = Failure(e) - def isThrow = true - def isReturn = false + def isThrow: Boolean = true + def isReturn: Boolean = false def throwable: Throwable = e - def rescue[R2 >: R](rescueException: PartialFunction[Throwable, Try[R2]]) = { + def rescue[R2 >: R](rescueException: PartialFunction[Throwable, Try[R2]]): Try[R2] = { try { val result = rescueException.applyOrElse(e, Throw.AlwaysNotApplied) if (result eq Throw.NotApplied) this else result @@ -251,16 +258,16 @@ final case class Throw[+R](e: Throwable) extends Try[R] { } } def apply(): R = throw e - def flatMap[R2](f: R => Try[R2]) = this.asInstanceOf[Throw[R2]] + def flatMap[R2](f: R => Try[R2]): Throw[R2] = this.asInstanceOf[Throw[R2]] def flatten[T](implicit ev: R <:< Try[T]): Try[T] = this.asInstanceOf[Throw[T]] def map[X](f: R => X): Try[X] = this.asInstanceOf[Throw[X]] def cast[X]: Try[X] = this.asInstanceOf[Throw[X]] - def exists(p: R => Boolean) = false - def filter(p: R => Boolean) = this - def withFilter(p: R => Boolean) = this - def onFailure(rescueException: Throwable => Unit) = { rescueException(e); this } - def onSuccess(f: R => Unit) = this - def handle[R2 >: R](rescueException: PartialFunction[Throwable, R2]) = + def exists(p: R => Boolean): Boolean = false + def filter(p: R => Boolean): Throw[R] = this + def withFilter(p: R => Boolean): Throw[R] = this + def onFailure(rescueException: Throwable => Unit): Throw[R] = { rescueException(e); this } + def onSuccess(f: R => Unit): Throw[R] = this + def handle[R2 >: R](rescueException: PartialFunction[Throwable, R2]): Try[R2] = if (rescueException.isDefinedAt(e)) { Try(rescueException(e)) } else { @@ -269,8 +276,8 @@ final case class Throw[+R](e: Throwable) extends Try[R] { } object Return { - val Unit = Return(()) - val Void = Return[Void](null) + val Unit: Return[Unit] = Return(()) + val Void: Return[Void] = Return[Void](null) val None: Return[Option[Nothing]] = Return(Option.empty) val Nil: Return[Seq[Nothing]] = Return(Seq.empty) val True: Return[Boolean] = Return(true) diff --git a/util-core/src/main/scala/com/twitter/util/TwitterDateFormat.scala b/util-core/src/main/scala/com/twitter/util/TwitterDateFormat.scala index b878e33130..6626557a3b 100644 --- a/util-core/src/main/scala/com/twitter/util/TwitterDateFormat.scala +++ b/util-core/src/main/scala/com/twitter/util/TwitterDateFormat.scala @@ -18,35 +18,44 @@ object TwitterDateFormat { new SimpleDateFormat(pattern, locale) } - def validatePattern(pattern: String): Unit = { - val stripped = stripSingleQuoted(pattern) - if (stripped.contains('Y') && !(stripped.contains('w'))) { + // Returns a quadruple of booleans, each signifying whether the pattern has + // a symbol for year-of-era (y), month-of-year (M), day-of-month (d), and + // hour-of-day (H), respectively. Throws an exception if week-year (Y) + // is used without week-in-year (w). + def validatePattern( + pattern: String, + newFormat: Boolean = false + ): (Boolean, Boolean, Boolean, Boolean) = { + var quoteOff = true + var optionalLevel = 0 + var has_Y = false + var has_w = false + var has_y = false + var has_M = false + var has_d = false + var has_H = false + val chars = pattern.toCharArray + var i = 0 + while (i < chars.length) { + chars(i) match { + case '\'' => quoteOff = !quoteOff + case '[' if newFormat && quoteOff => optionalLevel += 1 + case ']' if newFormat && quoteOff => optionalLevel -= 1 + case 'y' if !has_y && quoteOff && optionalLevel == 0 => has_y = true + case 'M' if !has_M && quoteOff && optionalLevel == 0 => has_M = true + case 'd' if !has_d && quoteOff && optionalLevel == 0 => has_d = true + case 'H' if !has_H && quoteOff && optionalLevel == 0 => has_H = true + case 'Y' if !has_Y && quoteOff && optionalLevel == 0 => has_Y = true + case 'w' if !has_w && quoteOff && optionalLevel == 0 => has_w = true + case _ => () + } + i += 1 + } + if (has_Y && !has_w) { throw new IllegalArgumentException( "Invalid date format uses 'Y' for week-of-year without 'w': %s".format(pattern) ) } - } - - /** Patterns can contain quoted strings 'foo' which we should ignore in checking the pattern */ - private[util] def stripSingleQuoted(pattern: String): String = { - val buf = new StringBuilder(pattern.size) - var startIndex = 0 - var endIndex = 0 - while (startIndex >= 0 && startIndex < pattern.size) { - endIndex = pattern.indexOf('\'', startIndex) - if (endIndex < 0) { - buf.append(pattern.substring(startIndex)) - startIndex = -1 - } else { - buf.append(pattern.substring(startIndex, endIndex)) - startIndex = endIndex + 1 - endIndex = pattern.indexOf('\'', startIndex) - if (endIndex < 0) { - throw new IllegalArgumentException("Unmatched quote in date format: %s".format(pattern)) - } - startIndex = endIndex + 1 - } - } - buf.toString + (has_y, has_M, has_d, has_H) } } diff --git a/util-core/src/main/scala/com/twitter/util/U64.scala b/util-core/src/main/scala/com/twitter/util/U64.scala deleted file mode 100644 index 19c91f8c44..0000000000 --- a/util-core/src/main/scala/com/twitter/util/U64.scala +++ /dev/null @@ -1,122 +0,0 @@ -// Copyright 2010 Twitter, Inc. -// -// Unsigned math on (64-bit) longs. - -package com.twitter.util - -import java.io.{ByteArrayInputStream, ByteArrayOutputStream, DataInputStream, DataOutputStream} -import scala.language.implicitConversions - -@deprecated("Use Java 8 unsigned Long APIs instead", "2018-02-13") -class RichU64Long(l64: Long) { - import U64._ - - def u64_<(y: Long): Boolean = u64_lt(l64, y) - - /** exclusive: (a, b) */ - def u64_within(a: Long, b: Long): Boolean = u64_lt(a, l64) && u64_lt(l64, b) - - /** inclusive: [a, b] */ - def u64_contained(a: Long, b: Long): Boolean = !u64_lt(l64, a) && !u64_lt(b, l64) - - def u64ToBigInt: BigInt = u64ToBigint(l64) - - def u64_/(y: Long): Long = - (u64ToBigInt / u64ToBigint(y)).longValue - - def u64_%(y: Long): Long = - (u64ToBigInt % u64ToBigint(y)).longValue - - // Other arithmetic operations don't need unsigned equivalents - // with 2s complement. - - def toU64ByteArray: Array[Byte] = { - val bytes = new ByteArrayOutputStream(8) - new DataOutputStream(bytes).writeLong(l64) - bytes.toByteArray - } - - def toU64HexString = (new RichU64ByteArray(toU64ByteArray)).toU64HexString -} - -@deprecated("Use Java 8 unsigned Long APIs instead", "2018-02-13") -class RichU64ByteArray(bytes: Array[Byte]) { - private def deSign(b: Byte): Int = { - if (b < 0) { - b + 256 - } else { - b - } - } - - // Note: this is padded, so we have lexical ordering. - def toU64HexString = bytes.map(deSign).map("%02x".format(_)).reduceLeft(_ + _) - def toU64Long = new DataInputStream(new ByteArrayInputStream(bytes)).readLong -} - -/** - * RichU64String parses string as a 64-bit unsigned hexadecimal long, - * outputting a signed Long or a byte array. Valid input is 1-16 hexadecimal - * digits (0-9, a-z, A-Z). - * - * @throws NumberFormatException if string is not a non-negative hexadecimal - * string - */ -@deprecated("Use Java 8 unsigned Long APIs instead", "2018-02-13") -class RichU64String(string: String) { - private[this] def validateHexDigit(c: Char): Unit = { - if (!(('0' <= c && c <= '9') || - ('a' <= c && c <= 'f') || - ('A' <= c && c <= 'F'))) { - throw new NumberFormatException("For input string: \"" + string + "\"") - } - } - - if (string.length > 16) { - throw new NumberFormatException("Number longer than 16 hex characters") - } else if (string.isEmpty) { - throw new NumberFormatException("Empty string") - } - string.foreach(validateHexDigit) - - def toU64ByteArray: Array[Byte] = { - val padded = "0" * (16 - string.length()) + string - (0 until 16 by 2) - .map(i => { - val parsed = Integer.parseInt(padded.slice(i, i + 2), 16) - assert(parsed >= 0) - parsed.toByte - }) - .toArray - } - - def toU64Long: Long = (new RichU64ByteArray(toU64ByteArray)).toU64Long -} - -@deprecated("Use Java 8 unsigned Long APIs instead", "2018-02-13") -object U64 { - private val bigInt0x8000000000000000L = (0x7FFFFFFFFFFFFFFFL: BigInt) + 1 - - val U64MAX = 0xFFFFFFFFFFFFFFFFL - val U64MIN = 0L - - def u64ToBigint(x: Long): BigInt = - if ((x & 0x8000000000000000L) != 0L) - ((x & 0x7FFFFFFFFFFFFFFFL): BigInt) + bigInt0x8000000000000000L - else - x: BigInt - - // compares x < y - def u64_lt(x: Long, y: Long): Boolean = - if (x < 0 == y < 0) - x < y // signed comparison, straightforward! - else - x > y // x is less if it doesn't have its high bit set (<0) - - implicit def longToRichU64Long(x: Long): RichU64Long = new RichU64Long(x) - - implicit def byteArrayToRichU64ByteArray(bytes: Array[Byte]): RichU64ByteArray = - new RichU64ByteArray(bytes) - - implicit def stringToRichU64String(string: String): RichU64String = new RichU64String(string) -} diff --git a/util-core/src/main/scala/com/twitter/util/Var.scala b/util-core/src/main/scala/com/twitter/util/Var.scala index a26ff4acdc..984cbd42ab 100644 --- a/util-core/src/main/scala/com/twitter/util/Var.scala +++ b/util-core/src/main/scala/com/twitter/util/Var.scala @@ -1,22 +1,53 @@ package com.twitter.util -import java.util.concurrent.atomic.{AtomicLong, AtomicReference, AtomicReferenceArray} import java.util.{List => JList} -import scala.annotation.tailrec -import scala.collection.JavaConverters._ -import scala.collection.generic.CanBuildFrom -import scala.collection.immutable +import scala.collection.{Seq => AnySeq} +import scala.collection.compat.immutable.ArraySeq +import scala.collection.immutable.SortedSet import scala.collection.mutable.Buffer +import scala.jdk.CollectionConverters._ import scala.language.higherKinds -import scala.reflect.ClassTag /** - * Trait Var represents a variable. It is a reference cell which is - * composable: dependent Vars (derived through flatMap) are - * recomputed automatically when independent variables change -- they - * implement a form of self-adjusting computation. + * Vars are values that vary over time. To create one, you must give it an + * initial value. * - * Vars are observed, notifying users whenever the variable changes. + * {{{ + * val a = Var[Int](1) + * }}} + * + * A Var created this way can be sampled to retrieve its current value, + * + * {{{ + * println(Var.sample(a)) // prints 1 + * }}} + * + * or, invoked to assign it new values. + * + * {{{ + * a.update(2) + * println(Var.sample(a)) // prints 2 + * }}} + * + * Vars can be derived from other Vars. + * + * {{{ + * val b = a.flatMap { x => Var(x + 2) } + * println(Var.sample(b)) // prints 4 + * }}} + * + * And, if the underlying is assigned a new value, the derived Var is updated. + * Updates are computed lazily, so while assignment is cheap, sampling is where + * we pay the cost of the computation required to derive the new Var. + * + * {{{ + * a.update(1) + * println(Var.sample(b)) // prints 3 + * }}} + * + * A key difference between the derived Var and its underlying is that derived + * Vars can't be assigned new values. That's why `b`, from the example above, + * can't be invoked to assign it a new value, it can only be sampled. * * @note Vars do not always perform the minimum amount of * re-computation. @@ -61,28 +92,52 @@ trait Var[+T] { self => * unobserved Var returned by flatMap will not invoke `f` */ def flatMap[U](f: T => Var[U]): Var[U] = new Var[U] { - def observe(depth: Int, obs: Observer[U]) = { - val inner = new AtomicReference(Closable.nop) - val outer = self.observe( - depth, - Observer(t => { - // TODO: Right now we rely on synchronous propagation; and - // thus also synchronous closes. We should instead perform - // asynchronous propagation so that it is is safe & - // predictable to have asynchronously closing Vars, for - // example. Currently the only source of potentially - // asynchronous closing is Var.async; here we have modified - // the external process to close asynchronously with the Var - // itself. Thus we know the code path here is synchronous: - // we control all Var implementations, and also all Closable - // combinators have been modified to evaluate their respective - // Futures eagerly. - val done = inner.getAndSet(f(t).observe(depth + 1, obs)).close() - assert(done.isDone) - }) - ) - - Closable.sequence(outer, Closable.ref(inner)) + private class ObserverManager(depth: Int, obs: Observer[U]) extends Closable with (T => Unit) { + private[this] var closable: Closable = Closable.nop + private[this] var closed: Boolean = false + + def apply(t: T): Unit = { + // We have to synchronize and make sure we're not already closed or else + // we may generate a new observation that has raced with the a `close` + // call and lost, therefore making an orphan and a resource leak. + val toClose = this.synchronized { + if (closed) Closable.nop + else { + val old = closable + closable = f(t).observe(depth + 1, obs) + old + } + } + // TODO: Right now we rely on synchronous propagation; and + // thus also synchronous closes. We should instead perform + // asynchronous propagation so that it is is safe & + // predictable to have asynchronously closing Vars, for + // example. Currently the only source of potentially + // asynchronous closing is Var.async; here we have modified + // the external process to close asynchronously with the Var + // itself. Thus we know the code path here is synchronous: + // we control all Var implementations, and also all Closable + // combinators have been modified to evaluate their respective + // Futures eagerly. + val done = toClose.close() + assert(done.isDefined) + } + + def close(deadline: Time): Future[Unit] = { + val toClose = this.synchronized { + closed = true + val old = closable + closable = Closable.nop + old + } + toClose.close(deadline) + } + } + + def observe(depth: Int, obs: Observer[U]): Closable = { + val manager = new ObserverManager(depth, obs) + val outer = self.observe(depth, Observer(manager)) + Closable.sequence(outer, manager) } } @@ -98,9 +153,7 @@ trait Var[+T] { self => */ lazy val changes: Event[T] = new Event[T] { def register(s: Witness[T]) = - self.observe { newv => - s.notify(newv) - } + self.observe { newv => s.notify(newv) } } /** @@ -127,7 +180,7 @@ object Var { * A Var observer. Observers are owned by exactly one producer, * enforced by a leasing mechanism. */ - private[util] class Observer[-T](observe: T => Unit) { + final class Observer[-T](observe: T => Unit) { private[this] var thisOwner: AnyRef = null private[this] var thisVersion = Long.MinValue @@ -158,8 +211,8 @@ object Var { } } - private[util] object Observer { - def apply[T](k: T => Unit) = new Observer(k) + object Observer { + def apply[T](k: T => Unit): Observer[T] = new Observer(k) } /** @@ -215,7 +268,7 @@ object Var { // In order to support unsubscribing from diffs when v is no longer referenced // we must avoid diffs keeping a strong reference to v. - val witness = Witness.weakReference { diff: Diff[CC, T] => + val witness = Witness.weakReference { (diff: Diff[CC, T]) => synchronized { v.update(diff.patch(v())) } @@ -226,18 +279,10 @@ object Var { v } - private case class Value[T](v: T) extends Var[T] { - protected def observe(depth: Int, obs: Observer[T]): Closable = { - obs.claim(this) - obs.publish(this, v, 0) - Closable.nop - } - } - /** * Create a new, constant, v-valued Var. */ - def value[T](v: T): Var[T] = Value(v) + def value[T](v: T): Var[T] with Extractable[T] = new ConstVar(v) /** * Collect a collection of Vars into a Var of collection. @@ -261,10 +306,8 @@ object Var { * // refCollectIndependent == Vector((1,2), (2,2), (2,4)) * }}} */ - def collect[T: ClassTag, CC[X] <: Traversable[X]]( - vars: CC[Var[T]] - )(implicit newBuilder: CanBuildFrom[CC[T], T, CC[T]]): Var[CC[T]] = { - val vs = vars.toArray + def collect[T](vars: AnySeq[Var[T]]): Var[Seq[T]] = { + val vs = vars.toIndexedSeq def tree(begin: Int, end: Int): Var[Seq[T]] = if (begin == end) Var(Seq.empty) @@ -278,16 +321,11 @@ object Var { } yield left ++ right } - tree(0, vs.length).map { ts => - val b = newBuilder() - b ++= ts - b.result() - } + tree(0, vs.length) } /** * Collect a collection of Vars into a Var of collection. - * Var.collect can result in a stack overflow if called with a large sequence. * Var.collectIndependent breaks composition with respect to update propagation. * That is, collectIndependent can fail to correctly update interdependent vars, * but is safe for independent vars. @@ -307,42 +345,50 @@ object Var { * // refCollectIndependent == Vector((1,2), (2,2), (2,4)) * }}} */ - def collectIndependent[T: ClassTag, CC[X] <: Traversable[X]]( - vars: CC[Var[T]] - )( - implicit newBuilder: CanBuildFrom[CC[T], T, CC[T]] - ): Var[CC[T]] = - async(newBuilder().result()) { v => + def collectIndependent[T](vars: AnySeq[Var[T]]): Var[Seq[T]] = + async(Seq.empty[T]) { u => val N = vars.size - val cur = new AtomicReferenceArray[T](N) - @volatile var filling = true - def build() = { - val b = newBuilder() - var i = 0 - while (i < N) { - b += cur.get(i) - i += 1 - } - b.result() - } - def publish(i: Int, newValue: T) = { - cur.set(i, newValue) - if (!filling) v() = build() + // `filling` represents whether or not we have gone through our collection + // of `Var`s and added observations. Once we have "filled" the `cur` array + // for the first time, we can publish an initial update to `u`. Subsequent, + // updates to each constintuent Var are guarded by a lock on the `cur` and + // we ensure that we publish the update before any new updates can come + // through. + // + // @note there is still a subtle race where we can be in the middle + // of "filling" and receive updates on previously filled Vars. However, this + // is an acceptable race since technically we haven't added observations to all + // of the constituent Vars yet. + // + // @note "filling" only works with the guarantee that the initial `observe` is + // synchronous. This should be the case with Vars since they have an initial value. + var filling = true + val cur = new Array[Any](N) + + def publishAt(i: Int): T => Unit = { newValue => + cur.synchronized { + cur(i) = newValue + if (filling && i == N - 1) filling = false + if (!filling) { + // toSeq does not do a deep copy until 2.13 + val copy = new Array[Any](N) + Array.copy(cur, 0, copy, 0, cur.length) + u() = ArraySeq.unsafeWrapArray(copy).asInstanceOf[Seq[T]] + } + } } - val closes = new Array[Closable](N) + val closables = new Array[Closable](N) var i = 0 - for (v <- vars) { - val j = i - closes(j) = v.observe(0, Observer(newValue => publish(j, newValue))) + val iter = vars.iterator + while (iter.hasNext) { + val v = iter.next() + closables(i) = v.observe(publishAt(i)) i += 1 } - filling = false - v() = build() - - Closable.all(closes: _*) + Closable.all(ArraySeq.unsafeWrapArray(closables): _*) } /** @@ -428,23 +474,15 @@ object Var { private object UpdatableVar { import Var.Observer - case class Party[T](obs: Observer[T], depth: Int, n: Long) { - @volatile var active = true + private final class Party[T](val observer: Observer[T], val depth: Int, val sequence: Long) { + @volatile var active: Boolean = true } - case class State[T](value: T, version: Long, parties: immutable.SortedSet[Party[T]]) { - def -(p: Party[T]) = copy(parties = parties - p) - def +(p: Party[T]) = copy(parties = parties + p) - def :=(newv: T) = copy(value = newv, version = version + 1) - } - - implicit def order[T] = new Ordering[Party[T]] { - // This is safe because observers are compared - // only from the same counter. - def compare(a: Party[T], b: Party[T]): Int = { - val c1 = a.depth compare b.depth + private val partyOrder: Ordering[Party[_]] = new Ordering[Party[_]] { + def compare(a: Party[_], b: Party[_]): Int = { + val c1 = java.lang.Integer.compare(a.depth, b.depth) if (c1 != 0) return c1 - a.n compare b.n + java.lang.Long.compare(a.sequence, b.sequence) } } } @@ -453,45 +491,70 @@ private[util] class UpdatableVar[T](init: T) extends Var[T] with Updatable[T] wi import UpdatableVar._ import Var.Observer - private[this] val n = new AtomicLong(0) - private[this] val state = new AtomicReference(State[T](init, 0, immutable.SortedSet.empty)) - - @tailrec - private[this] def cas(next: State[T] => State[T]): State[T] = { - val from = state.get - val to = next(from) - if (state.compareAndSet(from, to)) to else cas(next) - } - - def apply(): T = state.get.value + // The state must be mutated or read under a synchronized block on 'this'. The value can be read + // w/o an external synchronization. + @volatile private[this] var value = init + private[this] var version = 0L + private[this] var partySequence = 0L + @volatile private[this] var parties = + SortedSet.empty[Party[T]](partyOrder.asInstanceOf[Ordering[Party[T]]]) + + def apply(): T = value + + def update(newValue: T): Unit = { + val v = synchronized { + version += 1 + value = newValue + version + } - def update(newv: T): Unit = synchronized { - val State(value, version, parties) = cas(_ := newv) - for (p @ Party(obs, _, _) <- parties) { + // Another party maybe be racing to register while we're updating with the new value. + // This is not a big deal b/c observers prevent double-updates by design (via versions). + parties.foreach { p => // An antecedent update may have closed the current // party (e.g. flatMap does this); we need to check that // the party is active here in order to prevent stale updates. - if (p.active) - obs.publish(this, value, version) + if (p.active) p.observer.publish(this, newValue, v) } } protected def observe(depth: Int, obs: Observer[T]): Closable = { obs.claim(this) - val party = Party(obs, depth, n.getAndIncrement()) - val State(value, version, _) = cas(_ + party) - obs.publish(this, value, version) + + val (p, curValue, curVersion) = synchronized { + val party = new Party(obs, depth, partySequence) + parties = parties + party + partySequence += 1 + (party, value, version) + } + + obs.publish(this, curValue, curVersion) new Closable { - def close(deadline: Time) = { - party.active = false - cas(_ - party) + def close(deadline: Time): Future[Unit] = { + p.active = false + UpdatableVar.this.synchronized { + parties = parties - p + } Future.Done } } } - override def toString = "Var(" + state.get.value + ")@" + hashCode + override def toString: String = "Var(" + value + ")@" + hashCode +} + +/** + * A constant [[Extractable]] [[Var]] on `v`. + */ +class ConstVar[T](v: T) extends Var[T] with Extractable[T] { + protected def observe(depth: Int, obs: Var.Observer[T]): Closable = { + obs.claim(this) + obs.publish(this, v, 0) + Closable.nop + } + + def apply(): T = v } /** diff --git a/util-core/src/main/scala/com/twitter/util/WindowedAdder.scala b/util-core/src/main/scala/com/twitter/util/WindowedAdder.scala index e1ceffab5d..7088a3cc95 100644 --- a/util-core/src/main/scala/com/twitter/util/WindowedAdder.scala +++ b/util-core/src/main/scala/com/twitter/util/WindowedAdder.scala @@ -27,7 +27,10 @@ private[twitter] object WindowedAdder { } } -private[twitter] class WindowedAdder private[WindowedAdder] (window: Long, N: Int, now: () => Long) { +private[twitter] class WindowedAdder private[WindowedAdder] ( + window: Long, + N: Int, + now: () => Long) { private[this] val writer = new LongAdder() @volatile private[this] var gen = 0 private[this] val expiredGen = new AtomicInteger(gen) diff --git a/util-core/src/main/scala/com/twitter/util/WorkQueueFiber.scala b/util-core/src/main/scala/com/twitter/util/WorkQueueFiber.scala new file mode 100644 index 0000000000..9f5df55569 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/WorkQueueFiber.scala @@ -0,0 +1,271 @@ +package com.twitter.util + +import java.util.ArrayDeque +import java.util.concurrent.ConcurrentLinkedQueue +import java.util.concurrent.atomic.AtomicBoolean +import scala.annotation.tailrec + +/** + * An experimental [[Fiber]] implementation that has its own internal work queue + * + * This `Fiber` abstraction is intended to allow work associated with the fiber to have + * a common execution context while not defining the specific thread model that will execute + * this `Fiber`s work queue. + * + * The work queue is actually two work queues: one that accepts tasks submitted by this thread + * and one that is thread safe. This is to optimize for the common case where work can immediate + * finish other work and immediately trigger another task. This lets us avoid the overhead + * of using a thread safe queue in a very common case. + * + * @note that priority is given to the work submitted by this thread. + * + * @note care should be taken with the FiberTasks submitted to this Fiber. If a FiberTask + * should throw an exception, the fiber could get stuck in an `executing` state and + * resource tracking would be incomplete. + * + * @param fiberMetrics telemetry used for tracking the execution of this fiber, and likely + * other fibers as well. + */ +private[twitter] abstract class WorkQueueFiber(fiberMetrics: WorkQueueFiber.FiberMetrics) + extends Fiber { + + fiberMetrics.fiberCreated() + + // The thread safe queue that we can submit work to + private[this] val multiThreadedQueue = new ConcurrentLinkedQueue[FiberTask] + + // Flag used to ensure that at most one thread is processing this fibers work queue + // at any given time. + private[this] val executing = new AtomicBoolean(false) + + // State used for the non-thread safe queue. This is a micro opt to make + // it so we can quickly trampoline tasks, potentially without allocating. + private[this] var r0, r1, r2: FiberTask = null + private[this] val rs = new ArrayDeque[FiberTask](0 /* effectively 8 due to impl */ ) + + // This is thread safe. Why? + // * We only care if the value stored is *this* thread. + // * A thread will only ever set the field to itself when beginning execution + // and unsets the field on exiting the submission loop. + // + // So if we see a value that is equal to `Thread.currentThread`, it must + // be *this* thread that set it, and if it is *this* thread that set it, the access has been + // single threaded all along. + private[this] var executingThread: Thread = null + + private[this] val workRunnable: Runnable = { () => startWork() } + + private[twitter] val tracker: ResourceTracker = new ResourceTracker + + /** + * Submit the work to an actual thread for execution + * + * @note that we require this method call to be thread safe such that the memory updates + * from the thread that calls this method will be visible to the thread that begins + * the execution of `r`. This is generally true of thread-safe execution of `Runnable`s. + */ + protected def schedulerSubmit(r: Runnable): Unit + + final def submitTask(r: FiberTask): Unit = { + fiberMetrics.taskSubmissionIncrement() + // Check if anybody is running this already before getting a reference to the current thread. + if ((executingThread ne null) && (executingThread eq Thread.currentThread())) { + fiberMetrics.threadLocalSubmitIncrement() + threadLocalSubmit(r) + } else multiThreadedSubmit(r) + } + + /** + * If necessary, execute any scheduler specific flushing. + */ + protected def schedulerFlush(): Unit + + /** + * If necessary, drain any thread specific state and resubmit the work queue. + * + * @note This is in service of blocking calls which are not desirable in async code + * but sometimes you gotta do what you gotta do. It should be noted that in the + * absence of `flush()` calls this fiber can guarantee that tasks are executed + * serially, eg one task is fully completed before the next is begun. However, + * a `flush()` call will necessarily shift the work potentially to a separate + * thread allowing all the code after a `flush()` call to potentially run parallel + * with other work submitted to the fiber. + */ + final def flush(): Unit = { + val inFiber = Thread.currentThread() eq executingThread + // Record flush start time so we can account for the extra work done by this `flush` call. + val flushStartTime = if (inFiber) ResourceTracker.currentThreadCpuTime else 0l + if (inFiber) { + fiberMetrics.flushIncrement() + // We need to move all our work into a new `schedulerSubmit` call. That means moving all + // our thread local work into the thread-safe queue and resubmitting this fiber if there + // are still tasks queued. + while (localHasNext) { + multiThreadedQueue.add(localNext()) + } + + // Set this thread to non-executing and then try to re-submit ourselves, if necessary. + // Someone else might re-submit first, which is fine. + // + // Note: that we _must_ de-schedule ourselves in the case that the work we're waiting + // on belongs to this fiber but is completed later, eg, it wasn't already scheduled. + executingThread = null + executing.set(false) + + if (!multiThreadedQueue.isEmpty) { + // There is fiber work that needs to be executed. Attempt to reschedule. + tryScheduleFiber() + } + } + // Depending on the underlying execution engine, we may need to flush the scheduler as well. + schedulerFlush() + // Subtract the `flush` time from this fiber's total cpu time because the tasks executed may + // not belong to this fiber, and if they do, they'd be double counted since we should be in the + // `startWork` method right now executing one of the FiberTasks. + if (inFiber) { + val flushTime = ResourceTracker.currentThreadCpuTime - flushStartTime + if (flushTime > 0) tracker.addCpuTime(-flushTime) + } + } + + private[this] def threadLocalSubmit(r: FiberTask): Unit = { + if (r0 == null) r0 = r + else if (r1 == null) r1 = r + else if (r2 == null) r2 = r + else rs.addLast(r) + } + + private[this] def multiThreadedSubmit(r: FiberTask): Unit = { + multiThreadedQueue.add(r) + tryScheduleFiber() + } + + private[this] def tryScheduleFiber(): Unit = { + if (!executing.getAndSet(true)) { + // Lucky us! We're the ones that submit to the scheduler. + fiberMetrics.schedulerSubmissionIncrement() + schedulerSubmit(workRunnable) + } + } + + /** Start executing the work queue. + * + * Invoking this will execute all the tasks in both the local and multi-threaded work queues. + */ + private[this] def startWork(): Unit = { + + // This check serves two purposes. First, it guarantees that all updates to the non-thread + // safe state are visible to this thread because every time we exit the `startWork` loop we + // end by setting this state to false, and this thread reading it is a happens-before. We + // should also get that through the `schedulerSubmit` calls but this leaves no doubt. + // If it shows up as a performance bottleneck then we can consider removing it. + // + // Second, failure to see the expected state represents a very serious error: if executing + // is not true then something has gone catastrophically wrong. + if (!executing.get) { + throw new IllegalStateException("Fiber entered startWork loop with execution flag false") + } + + val currentThread = Thread.currentThread() + executingThread = currentThread + + // A thread just starting work should find two things: that the thread local queue is empty + // and that the multi-threaded queue is non-empty. + assert(!localHasNext) + assert(!multiThreadedQueue.isEmpty) + + // Expects that the first thing to do is execute a thread local task. + @tailrec def go(count: Int, n: FiberTask): Int = { + n.doRun() + // If we're not the executing thread anymore it is because a task has "flushed" us. + // If we've been flushed then we just need to exit now. The `flush()` method will + // take care of the state management and account for resource tracking. + if (executingThread ne currentThread) { + count + } else { + // Get our next task, if it exists. + var nn = localNext() + if (nn == null) { + nn = multiThreadedQueue.poll() + } + + if (nn != null) go(count + 1, nn) + else { + // We found our local queue and multi-threaded queues both empty. We now go about + // exiting this fiber submission. + + // We set `executingThread` to null before unsetting our executing flag to ensure + // that it is always null when entering this method. Otherwise we might race between + // unsetting the flag, thread pause, another thread submits to the scheduler, + // and finally this thread gets around to unsetting. + executingThread = null + executing.set(false) + + // Now we need to double check that someone didn't submit a task and not submit + // it to the scheduler because this loop was running. + if (!multiThreadedQueue.isEmpty && !executing.getAndSet(true)) { + executingThread = currentThread + go(count + 1, multiThreadedQueue.poll()) + } else { + count + } + } + } + } + + // Start tracking cpu usage + val trackingStartTime = ResourceTracker.currentThreadCpuTime + val continuationCount = go(1, multiThreadedQueue.poll()) + tracker.addContinuations(continuationCount) + // See how much CPU time has elapsed from this fiber submission. + val elapsedTime = ResourceTracker.currentThreadCpuTime - trackingStartTime + if (elapsedTime > 0) { + tracker.addCpuTime(elapsedTime) + } + } + + @inline private[this] def localHasNext: Boolean = r0 != null + + @inline private[this] def localNext(): FiberTask = { + // via moderately silly benchmarking, the + // queue unrolling gives us a ~50% speedup + // over pure Queue usage for common + // situations. + + val r = r0 + r0 = r1 + r1 = r2 + if (r2 != null) { + // If r2 was null then rs must be empty and no need to consult the + // heavy(ish) data structure. + r2 = rs.poll() + } + r + } +} + +private[twitter] object WorkQueueFiber { + + /** Telemetry used to monitor [[WorkQueueFiber]]s */ + abstract class FiberMetrics { + + /** Called when a new [[WorkQueueFiber]] is created */ + def fiberCreated(): Unit + + /** Called when a new task is submitted to a [[WorkQueueFiber]] via the thread-local queue */ + def threadLocalSubmitIncrement(): Unit + + /** Called when a task is submitted to a [[WorkQueueFiber]] */ + def taskSubmissionIncrement(): Unit + + /** Called when a [[WorkQueueFiber]] is submitted for execution */ + def schedulerSubmissionIncrement(): Unit + + /** + * Called when a [[Fiber]] performs a flush + * + * @note not invoked if the flush() call is a no-op + */ + def flushIncrement(): Unit + } +} diff --git a/util-core/src/main/scala/com/twitter/util/WrappedValue.scala b/util-core/src/main/scala/com/twitter/util/WrappedValue.scala new file mode 100644 index 0000000000..95feac3e99 --- /dev/null +++ b/util-core/src/main/scala/com/twitter/util/WrappedValue.scala @@ -0,0 +1,28 @@ +package com.twitter.util + +/** + * [[WrappedValue]] is a marker interface intended for use by case classes wrapping a single value. + * The intent is to signal that the underlying value should be directly serialized + * or deserialized ignoring the wrapping case class. + * + * [[WrappedValue]] is a [[https://docs.scala-lang.org/overviews/core/value-classes.html Universal Trait]] + * so that it can be used by [[https://docs.scala-lang.org/overviews/core/value-classes.html value classes]]. + * + * ==Usage== + * {{{ + * case class WrapperA(underlying: Int) extends WrappedValue[Int] + * case class WrapperB(underlying: Int) extends AnyVal with WrappedValue[Int] + * }}} + */ +trait WrappedValue[T] extends Any { self: Product => + + def onlyValue: T = self.productElement(0).asInstanceOf[T] + + def asString: String = onlyValue.toString + + /** + * @note we'd rather not override `toString`, but this is required to correctly handle + * [[WrappedValue]] instances used as Map keys. + */ + override def toString: String = asString +} diff --git a/util-core/src/main/scala/com/twitter/util/package.scala b/util-core/src/main/scala/com/twitter/util/package.scala deleted file mode 100644 index 4adb0b95ab..0000000000 --- a/util-core/src/main/scala/com/twitter/util/package.scala +++ /dev/null @@ -1,23 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter - -package object util { - // backward compat for people who imported "com.twitter.util.TimeConversions._" - val TimeConversions = com.twitter.conversions.time - val StorageUnitConversions = com.twitter.conversions.storage -} diff --git a/util-core/src/test/java/BUILD b/util-core/src/test/java/BUILD index 05f296d36e..4d9e1bb8aa 100644 --- a/util-core/src/test/java/BUILD +++ b/util-core/src/test/java/BUILD @@ -1,14 +1,10 @@ -junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scala-lang:scala-library", - "3rdparty/jvm/org/scalacheck", - "util/util-core/src/main/java", - "util/util-core/src/main/scala", - "util/util-function/src/main/java", + "util/util-core/src/test/java/com/twitter/concurrent", + "util/util-core/src/test/java/com/twitter/io", + "util/util-core/src/test/java/com/twitter/service", + "util/util-core/src/test/java/com/twitter/util", + "util/util-core/src/test/java/com/twitter/util/javainterop", ], ) diff --git a/util-core/src/test/java/com/twitter/concurrent/BUILD b/util-core/src/test/java/com/twitter/concurrent/BUILD new file mode 100644 index 0000000000..7258f00909 --- /dev/null +++ b/util-core/src/test/java/com/twitter/concurrent/BUILD @@ -0,0 +1,17 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core/src/main/java/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/concurrent", + ], +) diff --git a/util-core/src/test/java/com/twitter/concurrent/SpoolCompilationTest.java b/util-core/src/test/java/com/twitter/concurrent/SpoolCompilationTest.java index 19263fb872..86ee3e5fcd 100644 --- a/util-core/src/test/java/com/twitter/concurrent/SpoolCompilationTest.java +++ b/util-core/src/test/java/com/twitter/concurrent/SpoolCompilationTest.java @@ -4,35 +4,12 @@ import com.twitter.util.Future; import org.junit.Assert; import org.junit.Test; -import scala.collection.JavaConversions; +import scala.collection.JavaConverters; import java.util.Arrays; import java.util.Collection; public class SpoolCompilationTest { - private static class OwnSpool extends AbstractSpool { - @Override - public boolean isEmpty() { - return false; - } - - @Override - public Future> tail() { - return Future.value(Spools.newEmptySpool()); - } - - @Override - public String head() { - return "spool"; - } - } - - @Test - public void testOwnSpool() { - Spool a = new OwnSpool(); - Assert.assertFalse(a.isEmpty()); - Assert.assertEquals("spool", a.head()); - } @Test public void testSpoolCreation() { @@ -49,18 +26,18 @@ public void testSpoolCreation() { public void testSpoolConcat() throws Exception { Spool a = Spools.newSpool(Arrays.asList("a")); Spool b = Spools.newSpool(Arrays.asList("b")); - Spool c = Spools.newSpool(Arrays.asList("c")); + Spool cd = Spools.newSpool(Arrays.asList("c", "d")); Spool ab = a.concat(b); Spool abNothing = ab.concat(Spools.newEmptySpool()); - Spool abc = Await.result(ab.concat(Future.value(c))); + Spool abcd = Await.result(ab.concat(Future.value(cd))); - Collection listA = JavaConversions.seqAsJavaList(Await.result(ab.toSeq())); - Collection listB = JavaConversions.seqAsJavaList(Await.result(abNothing.toSeq())); - Collection listC = JavaConversions.seqAsJavaList(Await.result(abc.toSeq())); + Collection listA = JavaConverters.seqAsJavaListConverter(Await.result(ab.toSeq())).asJava(); + Collection listB = JavaConverters.seqAsJavaListConverter(Await.result(abNothing.toSeq())).asJava(); + Collection listC = JavaConverters.seqAsJavaListConverter(Await.result(abcd.toSeq())).asJava(); Assert.assertArrayEquals(new String[] { "a", "b"}, listA.toArray()); Assert.assertArrayEquals(new String[] { "a", "b"}, listB.toArray()); - Assert.assertArrayEquals(new String[] { "a", "b", "c"}, listC.toArray()); + Assert.assertArrayEquals(new String[] { "a", "b", "c" , "d"}, listC.toArray()); } } diff --git a/util-core/src/test/java/com/twitter/io/BUILD b/util-core/src/test/java/com/twitter/io/BUILD new file mode 100644 index 0000000000..596b31d145 --- /dev/null +++ b/util-core/src/test/java/com/twitter/io/BUILD @@ -0,0 +1,17 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:scala", + "util/util-core/src/main/java/com/twitter/io", + ], +) diff --git a/util-core/src/test/java/com/twitter/io/BufCompilationTest.java b/util-core/src/test/java/com/twitter/io/BufCompilationTest.java index d1e98587c8..386ae74b9c 100644 --- a/util-core/src/test/java/com/twitter/io/BufCompilationTest.java +++ b/util-core/src/test/java/com/twitter/io/BufCompilationTest.java @@ -3,7 +3,7 @@ import org.junit.Assert; import org.junit.Test; import scala.Option; -import scala.collection.JavaConversions; +import scala.collection.JavaConverters; import java.nio.ByteBuffer; import java.util.Arrays; @@ -55,7 +55,8 @@ public void testBufApply() { Buf abc = Bufs.utf8Buf("abc"); Buf def = Bufs.utf8Buf("def"); Buf merged = Buf.apply( - JavaConversions.asScalaBuffer(Arrays.asList(empty, abc, def))); + JavaConverters.asScalaBufferConverter(Arrays.asList(empty, abc, def)).asScala() + ); } @Test diff --git a/util-core/src/test/java/com/twitter/io/ReaderCompilationTest.java b/util-core/src/test/java/com/twitter/io/ReaderCompilationTest.java deleted file mode 100644 index 2a316c76e6..0000000000 --- a/util-core/src/test/java/com/twitter/io/ReaderCompilationTest.java +++ /dev/null @@ -1,55 +0,0 @@ -/* Copyright 2015 Twitter, Inc. */ -package com.twitter.io; - -import com.twitter.concurrent.AsyncStream; -import java.io.ByteArrayInputStream; -import java.io.File; - -import org.junit.Test; - -import static org.junit.Assert.assertNotNull; - -public class ReaderCompilationTest { - - @Test - public void testNull() { - assertNotNull(Readers.NULL); - } - - @Test - public void testNewBufReader() { - Readers.newBufReader(Bufs.EMPTY); - } - - @Test - public void testReadAll() { - Readers.readAll(Readers.newEmptyReader()); - } - - @Test - public void testConcat() { - Readers.concat(AsyncStream.>empty()); - } - - - @Test - public void testWritable() { - Pipe w = Readers.writable(); - w.close(); - } - - @Test - public void testNewFileReader() { - try { - Readers.newFileReader(new File("a")); - } catch (Exception x) { - // ok - } - } - - @Test - public void testNewInputStreamReader() throws Exception { - Readers.newInputStreamReader(new ByteArrayInputStream(new byte[0])); - } - -} diff --git a/util-core/src/test/java/com/twitter/io/WriterCompilationTest.java b/util-core/src/test/java/com/twitter/io/WriterCompilationTest.java deleted file mode 100644 index ce8502f98a..0000000000 --- a/util-core/src/test/java/com/twitter/io/WriterCompilationTest.java +++ /dev/null @@ -1,16 +0,0 @@ -/* Copyright 2015 Twitter, Inc. */ -package com.twitter.io; - -import java.io.ByteArrayOutputStream; - -import org.junit.Test; - -public class WriterCompilationTest { - - @Test - public void testNewOutputStreamWriter() { - Writer.ClosableWriter w = Writers.newOutputStreamWriter(new ByteArrayOutputStream()); - w.close(); - } - -} diff --git a/util-core/src/test/java/com/twitter/service/BUILD b/util-core/src/test/java/com/twitter/service/BUILD new file mode 100644 index 0000000000..b0a9acc072 --- /dev/null +++ b/util-core/src/test/java/com/twitter/service/BUILD @@ -0,0 +1,15 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core/src/main/java/com/twitter/service", + ], +) diff --git a/util-core/src/test/java/com/twitter/service/TestVirtualMachineErrorTerminator.java b/util-core/src/test/java/com/twitter/service/TestVirtualMachineErrorTerminator.java index 15e5cb1369..819daafd68 100644 --- a/util-core/src/test/java/com/twitter/service/TestVirtualMachineErrorTerminator.java +++ b/util-core/src/test/java/com/twitter/service/TestVirtualMachineErrorTerminator.java @@ -1,11 +1,14 @@ package com.twitter.service; -import junit.framework.TestCase; +import org.junit.After; +import org.junit.Before; +import org.junit.Test; +import static org.junit.Assert.*; /** * @author Attila Szegedi */ -public class TestVirtualMachineErrorTerminator extends TestCase +public class TestVirtualMachineErrorTerminator { private static boolean exitInvoked; private static final Object lock = new Object(); @@ -21,10 +24,12 @@ public void run() { }); } - protected void tearDown() { + @After + public void tearDown() { VirtualMachineErrorTerminator.reset(); } + @Test public void testNotTerminatingWhenNew() { assertNotTerminating(); } @@ -38,16 +43,19 @@ private static void assertNotTerminating() { } } + @Test public void testNotTerminatingOnNull() { VirtualMachineErrorTerminator.terminateVMIfMust(null); assertNotTerminating(); } + @Test public void testNotTerminatingOnOtherException() { VirtualMachineErrorTerminator.terminateVMIfMust(new Exception()); assertNotTerminating(); } + @Test public void testTerminatingOnNestedVirtualMachineError() throws Exception { Exception e = new Exception(); OutOfMemoryError oome = new OutOfMemoryError(); @@ -56,11 +64,13 @@ public void testTerminatingOnNestedVirtualMachineError() throws Exception { assertTerminated(); } + @Test public void testTerminatingOnVirtualMachineError() throws Exception { VirtualMachineErrorTerminator.terminateVMIfMust(new OutOfMemoryError()); assertTerminated(); } + @Test public void testTerminateDoesntAcceptNull() throws Exception { try { VirtualMachineErrorTerminator.terminateVM(null); @@ -70,6 +80,7 @@ public void testTerminateDoesntAcceptNull() throws Exception { } } + @Test public void testSecondTerminateIneffective() throws Exception { OutOfMemoryError oome1 = new OutOfMemoryError(); OutOfMemoryError oome2 = new OutOfMemoryError(); diff --git a/util-core/src/test/java/com/twitter/util/AwaitableCompilationTest.java b/util-core/src/test/java/com/twitter/util/AwaitableCompilationTest.java index 3fdae8ea17..5182dec164 100644 --- a/util-core/src/test/java/com/twitter/util/AwaitableCompilationTest.java +++ b/util-core/src/test/java/com/twitter/util/AwaitableCompilationTest.java @@ -12,17 +12,17 @@ public class AwaitableCompilationTest { public static class MockAwaitable implements Awaitable { @Override - public Awaitable ready(Duration timeout, CanAwait permit) throws TimeoutException { + public Awaitable ready(Duration timeout, Awaitable.CanAwait permit) throws TimeoutException { return this; } @Override - public String result(Duration timeout, CanAwait permit) throws Exception { + public String result(Duration timeout, Awaitable.CanAwait permit) throws Exception { return "42"; } @Override - public boolean isReady(CanAwait permit) { + public boolean isReady(Awaitable.CanAwait permit) { return true; } } diff --git a/util-core/src/test/java/com/twitter/util/BUILD b/util-core/src/test/java/com/twitter/util/BUILD new file mode 100644 index 0000000000..ac362447c8 --- /dev/null +++ b/util-core/src/test/java/com/twitter/util/BUILD @@ -0,0 +1,35 @@ +java_library( + name = "object_size_calculator", + sources = ["ObjectSizeCalculator.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/com/google/errorprone:error_prone_annotations", + ], +) + +junit_tests( + sources = [ + "!ObjectSizeCalculator.java", + "*.java", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + ":object_size_calculator", + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-core/src/main/java/com/twitter/util", + ], +) diff --git a/util-core/src/test/java/com/twitter/util/ObjectSizeCalculator.java b/util-core/src/test/java/com/twitter/util/ObjectSizeCalculator.java index 86efbb72d1..79fe00b97f 100644 --- a/util-core/src/test/java/com/twitter/util/ObjectSizeCalculator.java +++ b/util-core/src/test/java/com/twitter/util/ObjectSizeCalculator.java @@ -1,7 +1,8 @@ package com.twitter.util; +import com.sun.management.HotSpotDiagnosticMXBean; +import com.sun.management.VMOption; import java.lang.management.ManagementFactory; -import java.lang.management.MemoryPoolMXBean; import java.lang.reflect.Array; import java.lang.reflect.Field; import java.lang.reflect.Modifier; @@ -357,54 +358,51 @@ static MemoryLayoutSpecification getEffectiveMemoryLayoutSpecification() { dataModel + "' of sun.arch.data.model system property"); } - final String strVmVersion = System.getProperty("java.vm.version"); - final int vmVersion = Integer.parseInt(strVmVersion.substring(0, - strVmVersion.indexOf('.'))); - if (vmVersion >= 17) { - long maxMemory = 0; - for (MemoryPoolMXBean mp : ManagementFactory.getMemoryPoolMXBeans()) { - maxMemory += mp.getUsage().getMax(); - } - if (maxMemory < 30L * 1024 * 1024 * 1024) { - // HotSpot 17.0 and above use compressed OOPs below 30GB of RAM total - // for all memory pools (yes, including code cache). - return new MemoryLayoutSpecification() { - @Override public int getArrayHeaderSize() { - return 16; - } - @Override public int getObjectHeaderSize() { - return 12; - } - @Override public int getObjectPadding() { - return 8; - } - @Override public int getReferenceSize() { - return 4; - } - @Override public int getSuperclassFieldPadding() { - return 4; - } - }; - } + final HotSpotDiagnosticMXBean diagBean = ManagementFactory.getPlatformMXBean(HotSpotDiagnosticMXBean.class); + if (diagBean == null) { + throw new UnsupportedOperationException("ObjectSizeCalculator only supported on HotSpot VM"); } + final VMOption useCoopsOption = diagBean.getVMOption("UseCompressedOops"); + final boolean useCoops = Boolean.parseBoolean(useCoopsOption.getValue()); - // In other cases, it's a 64-bit uncompressed OOPs object model - return new MemoryLayoutSpecification() { - @Override public int getArrayHeaderSize() { - return 24; - } - @Override public int getObjectHeaderSize() { - return 16; - } - @Override public int getObjectPadding() { - return 8; - } - @Override public int getReferenceSize() { - return 8; - } - @Override public int getSuperclassFieldPadding() { - return 8; - } - }; + if (useCoops) { + // Returning the 64-bit compressed OOPs object model + return new MemoryLayoutSpecification() { + @Override public int getArrayHeaderSize() { + return 16; + } + @Override public int getObjectHeaderSize() { + return 12; + } + @Override public int getObjectPadding() { + return 8; + } + @Override public int getReferenceSize() { + return 4; + } + @Override public int getSuperclassFieldPadding() { + return 4; + } + }; + } else { + // Returning the 64-bit non-compressed OOPs object model + return new MemoryLayoutSpecification() { + @Override public int getArrayHeaderSize() { + return 24; + } + @Override public int getObjectHeaderSize() { + return 16; + } + @Override public int getObjectPadding() { + return 8; + } + @Override public int getReferenceSize() { + return 8; + } + @Override public int getSuperclassFieldPadding() { + return 8; + } + }; + } } } diff --git a/util-core/src/test/java/com/twitter/util/javainterop/BUILD b/util-core/src/test/java/com/twitter/util/javainterop/BUILD new file mode 100644 index 0000000000..8b74c78558 --- /dev/null +++ b/util-core/src/test/java/com/twitter/util/javainterop/BUILD @@ -0,0 +1,16 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core/src/main/java/com/twitter/util/javainterop", + ], +) diff --git a/util-core/src/test/resources/1.txt b/util-core/src/test/resources/1.txt new file mode 100644 index 0000000000..dc8344c1ce --- /dev/null +++ b/util-core/src/test/resources/1.txt @@ -0,0 +1 @@ +Lorem ipsum dolor sit amet diff --git a/util-core/src/test/resources/BUILD b/util-core/src/test/resources/BUILD new file mode 100644 index 0000000000..301b666c21 --- /dev/null +++ b/util-core/src/test/resources/BUILD @@ -0,0 +1,4 @@ +resources( + sources = ["**/*"], + tags = ["bazel-compatible"], +) diff --git a/util-core/src/test/resources/com/twitter/io/resource-test-file.txt b/util-core/src/test/resources/com/twitter/io/resource-test-file.txt new file mode 100644 index 0000000000..6be28d302a --- /dev/null +++ b/util-core/src/test/resources/com/twitter/io/resource-test-file.txt @@ -0,0 +1,23 @@ +Nam bibendum cursus risus, at volutpat elit tincidunt eu. Aliquam id ligula in mi tempor ultricies. +Phasellus quis consequat metus, at congue ipsum. Nam eu interdum elit, scelerisque laoreet neque. +Proin efficitur nisl lectus, in condimentum sem fringilla eget. Morbi condimentum non ipsum in +auctor. Mauris at turpis faucibus, sodales massa quis, commodo nisl. + +Aliquam hendrerit, ante eget accumsan lobortis, nibh nunc varius velit, ac vulputate sem odio id +erat. Pellentesque a imperdiet ipsum. Sed diam augue, eleifend ut congue at, accumsan sed libero. +Nulla ipsum dui, rhoncus a purus sit amet, sagittis convallis sapien. Curabitur auctor volutpat +laoreet. Donec at elit est. Aliquam ultrices urna fringilla sagittis iaculis. Sed sit amet +accumsan eros. Nullam tristique nibh sed molestie varius. Vestibulum quis dictum orci, et +imperdiet nunc. Maecenas aliquet, justo ut finibus malesuada, erat nunc pellentesque lacus, sit +amet iaculis libero dolor in libero. Interdum et malesuada fames ac ante ipsum primis in faucibus. +Quisque metus libero, pulvinar in sagittis eget, commodo eu justo. Morbi auctor tincidunt nibh vel +dictum. + +Vivamus ut elementum lectus. In vehicula eget lectus eget ultrices. Ut vitae est mattis, imperdiet +nulla non, molestie mi. Duis a ultrices nisl, eu consectetur neque. Integer id quam id massa +mattis rhoncus vitae eget ligula. Praesent dapibus, elit sit amet sollicitudin gravida, velit nisl +pellentesque ante, id tempor ante ligula quis est. Donec quis diam nec magna porttitor rhoncus. +Pellentesque pellentesque ullamcorper felis, malesuada fringilla odio accumsan eget. Praesent +ultrices turpis in libero convallis, a porta lectus venenatis. Curabitur aliquet, lorem eu cursus +scelerisque, lorem mi luctus mauris, in rhoncus leo neque sed odio. Nunc efficitur posuere dolor, +in aliquet orci congue vel. diff --git a/util-core/src/test/resources/empty-file.txt b/util-core/src/test/resources/empty-file.txt new file mode 100644 index 0000000000..e69de29bb2 diff --git a/util-core/src/test/resources/foo/bar/test-file.txt b/util-core/src/test/resources/foo/bar/test-file.txt new file mode 100644 index 0000000000..dddb718a87 --- /dev/null +++ b/util-core/src/test/resources/foo/bar/test-file.txt @@ -0,0 +1,3 @@ +Proin mattis nisi ac ipsum bibendum, eget lobortis lacus lobortis. Vivamus interdum ac tortor et accumsan. Aliquam scelerisque diam nibh, non efficitur dui eleifend luctus. Ut euismod risus a nunc tempus aliquet. Nulla facilisi. Vestibulum tempor justo id nulla rhoncus mattis. Praesent vitae congue diam. + +Sed euismod dictum fermentum. Nunc est dolor, auctor id vulputate id, commodo aliquam orci. Nam tincidunt vehicula sapien vitae dapibus. Proin nec metus magna. Suspendisse eget semper ex. Donec non mi lacus. Aliquam magna augue, eleifend quis scelerisque varius, congue et sapien. Praesent lacinia lorem quis urna efficitur, ut semper arcu blandit. diff --git a/util-core/src/test/resources/test.txt b/util-core/src/test/resources/test.txt new file mode 100644 index 0000000000..600d6e0166 --- /dev/null +++ b/util-core/src/test/resources/test.txt @@ -0,0 +1,21 @@ +Lorem ipsum dolor sit amet, consectetur adipiscing elit. Fusce et orci et felis aliquet luctus. In +a turpis tempus dolor condimentum pulvinar. Morbi pharetra suscipit velit. Phasellus et nisi +felis. Sed porta sapien vitae sem placerat, in vehicula massa cursus. Ut cursus ligula in justo +tincidunt pretium quis sit amet lorem. Phasellus feugiat, ante sed mattis sagittis, elit sapien +rhoncus quam, at varius enim eros faucibus metus. Suspendisse purus leo, convallis id luctus ut, +convallis nec neque. Suspendisse ut iaculis nisi. Vivamus cursus viverra faucibus. Vestibulum in +ligula at massa feugiat malesuada eget et urna. Sed nulla elit, rutrum et purus vel, aliquet +elementum risus. Suspendisse potenti. + +Vestibulum pellentesque, risus eget feugiat porttitor, mauris ligula bibendum risus, nec varius nisi +leo nec neque. Phasellus et sollicitudin elit, vitae ornare ante. Ut tempor risus rhoncus +tincidunt commodo. Phasellus efficitur nisl id lorem efficitur, vitae pharetra tortor fringilla. +Suspendisse tristique eros vel leo consectetur dictum. Praesent efficitur dolor ultricies, aliquam +odio et, consequat lorem. Proin dictum nunc sit amet nunc venenatis, eget venenatis nulla commodo. +Etiam ut mi sed sem interdum tincidunt. Phasellus magna tellus, suscipit sed commodo eu, +sollicitudin eget odio. Nulla massa elit, dignissim a quam ut, auctor dapibus augue. Etiam +interdum neque eu lectus ultricies, sed lacinia purus scelerisque. Proin nec consequat libero. Sed +eros nibh, cursus nec efficitur quis, pharetra ut nisl. Ut tincidunt, ex vel semper tempus, odio +leo pharetra eros, id interdum diam enim ut tortor. + +Sed ultricies luctus tempus. diff --git a/util-core/src/test/scala/BUILD b/util-core/src/test/scala/BUILD index 09d037242a..5cb247c531 100644 --- a/util-core/src/test/scala/BUILD +++ b/util-core/src/test/scala/BUILD @@ -1,20 +1,9 @@ -junit_tests( - sources = rglobs("*.scala"), - # fatal_warnings runs into an issue with scalatest/Future.isDone: - # - # [warn] ...: possible missing interpolator: detected interpolated identifier `$conforms` - # [warn] assert(future.close().isDone) - # [warn] ^ - # - # At some point later, when we can enable `-Xlint:-warn-missing-interpolator` - # for this target, we could try enabling this again. - fatal_warnings = False, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalacheck", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-core/src/test/java", + "util/util-core/src/test/scala/com/twitter/concurrent", + "util/util-core/src/test/scala/com/twitter/conversions", + "util/util-core/src/test/scala/com/twitter/io", + "util/util-core/src/test/scala/com/twitter/util", ], ) diff --git a/util-core/src/test/scala/com/twitter/concurrent/AsyncMeterTest.scala b/util-core/src/test/scala/com/twitter/concurrent/AsyncMeterTest.scala index e91e27c231..a05c6b310d 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/AsyncMeterTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/AsyncMeterTest.scala @@ -1,21 +1,37 @@ package com.twitter.concurrent import com.twitter.util._ -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import java.util.concurrent.{RejectedExecutionException, CancellationException} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class AsyncMeterTest extends FunSuite { +class AsyncMeterTest extends AnyFunSuite { import AsyncMeter._ + // Workaround methods for dealing with Scala compiler warnings: + // + // assert(f.isDone) cannot be called directly because it results in a Scala compiler + // warning: 'possible missing interpolator: detected interpolated identifier `$conforms`' + // + // This is due to the implicit evidence required for `Future.isDone` which checks to see + // whether the Future that is attempting to have `isDone` called on it conforms to the type + // of `Future[Unit]`. This is done using `Predef.$conforms` + // https://www.scala-lang.org/api/2.12.2/scala/Predef$.html#$conforms[A]:A%3C:%3CA + // + // Passing that directly to `assert` causes problems because the `$conforms` is also seen as + // an interpolated string. We get around it by evaluating first and passing the result to + // `assert`. + private[this] def isDone(f: Future[Unit]): Boolean = + f.isDefined + + private[this] def assertIsDone(f: Future[Unit]): Unit = + assert(isDone(f)) + test("AsyncMeter shouldn't wait at all when there aren't any waiters.") { val timer = new MockTimer val meter = newMeter(1, 1.second, 100)(timer) val result = meter.await(1) - assert(result.isDone) + assertIsDone(result) } test("AsyncMeter should allow more than one waiter and allow them on the schedule.") { @@ -24,12 +40,12 @@ class AsyncMeterTest extends FunSuite { val meter = newMeter(1, 1.second, 100)(timer) val ready = meter.await(1) val waiter = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) assert(!waiter.isDefined) ctl.advance(1.second) timer.tick() - assert(waiter.isDone) + assertIsDone(waiter) } } @@ -38,7 +54,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(1, 1.second, 100)(timer) val ready = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) val waiter = meter.await(1) assert(!waiter.isDefined) @@ -48,7 +64,7 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.second) timer.tick() - assert(waiter.isDone) + assertIsDone(waiter) } } @@ -57,7 +73,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(1, 1.second, 1)(timer) val ready = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) val waiter = meter.await(1) val rejected = meter.await(1) @@ -69,7 +85,7 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.second) timer.tick() - waiter.isDone + assertIsDone(waiter) } } @@ -80,7 +96,7 @@ class AsyncMeterTest extends FunSuite { var nr = 0 val ready = meter.await(1) val waiter = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) assert(!waiter.isDefined) val e = new Exception("boom!") @@ -98,7 +114,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(2, 1.second, 100)(timer) val ready = meter.await(2) - assert(ready.isDone) + assertIsDone(ready) val first = meter.await(1) val second = meter.await(1) @@ -107,8 +123,8 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.second) timer.tick() - assert(first.isDone) - assert(second.isDone) + assertIsDone(first) + assertIsDone(second) } } @@ -117,7 +133,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(1, 500.microseconds, 100)(timer) val ready = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) val first = meter.await(1) val second = meter.await(1) @@ -126,8 +142,8 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.millisecond) timer.tick() - assert(first.isDone) - assert(second.isDone) + assertIsDone(first) + assertIsDone(second) } } @@ -136,7 +152,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(2, 500.microseconds, 100)(timer) val ready = meter.await(2) - assert(ready.isDone) + assertIsDone(ready) val first = meter.await(1) val second = meter.await(2) @@ -147,9 +163,9 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.millisecond) timer.tick() - assert(first.isDone) - assert(second.isDone) - assert(third.isDone) + assertIsDone(first) + assertIsDone(second) + assertIsDone(third) } } @@ -175,7 +191,7 @@ class AsyncMeterTest extends FunSuite { Time.withCurrentTimeFrozen { ctl => val meter = newMeter(2, 1.second, 100)(timer) val ready = meter.await(2) - assert(ready.isDone) + assertIsDone(ready) val waiter = meter.await(2) assert(!waiter.isDefined) @@ -186,7 +202,7 @@ class AsyncMeterTest extends FunSuite { ctl.advance(500.milliseconds) timer.tick() - assert(waiter.isDone) + assertIsDone(waiter) } } @@ -206,7 +222,7 @@ class AsyncMeterTest extends FunSuite { val ready = meter.await(3) val first = meter.await(4) val second = meter.await(4) - assert(ready.isDone) + assertIsDone(ready) assert(!first.isDefined) assert(!second.isDefined) } @@ -219,16 +235,16 @@ class AsyncMeterTest extends FunSuite { val first = meter.await(1) val second = meter.await(1) val third = meter.await(1) - assert(ready.isDone) + assertIsDone(ready) ctl.advance(1.millisecond) timer.tick() - assert(first.isDone) + assertIsDone(first) assert(!second.isDefined) assert(!third.isDefined) ctl.advance(1.millisecond) timer.tick() - assert(second.isDone) - assert(third.isDone) + assertIsDone(second) + assertIsDone(third) } } @@ -240,7 +256,7 @@ class AsyncMeterTest extends FunSuite { assert(!greedy.isDefined) ctl.advance(1.second) timer.tick() - assert(greedy.isDone) + assertIsDone(greedy) } } @@ -257,7 +273,7 @@ class AsyncMeterTest extends FunSuite { } ctl.advance(1.second) timer.tick() - assert(first.isDone) + assertIsDone(first) } } @@ -280,7 +296,7 @@ class AsyncMeterTest extends FunSuite { ctl.advance(1.second) timer.tick() - assert(first.isDone) + assertIsDone(first) } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/AsyncMutexTest.scala b/util-core/src/test/scala/com/twitter/concurrent/AsyncMutexTest.scala index 81887c5bf3..d01281401e 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/AsyncMutexTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/AsyncMutexTest.scala @@ -1,13 +1,9 @@ package com.twitter.concurrent -import org.junit.runner.RunWith -import org.scalatest.FlatSpec -import org.scalatest.junit.JUnitRunner - import com.twitter.util.Await +import org.scalatest.flatspec.AnyFlatSpec -@RunWith(classOf[JUnitRunner]) -class AsyncMutexTest extends FlatSpec { +class AsyncMutexTest extends AnyFlatSpec { "AsyncMutex" should "admit only one operation at a time" in { val m = new AsyncMutex diff --git a/util-core/src/test/scala/com/twitter/concurrent/AsyncQueueTest.scala b/util-core/src/test/scala/com/twitter/concurrent/AsyncQueueTest.scala index 66c6dbd60d..38658bb7f1 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/AsyncQueueTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/AsyncQueueTest.scala @@ -1,13 +1,13 @@ package com.twitter.concurrent -import com.twitter.util.{Await, Return, Throw} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.util.{Await, Awaitable, Duration, Return, Throw} import scala.collection.immutable.Queue +import org.scalatest.funsuite.AnyFunSuite + +class AsyncQueueTest extends AnyFunSuite { + + private def await[T](t: Awaitable[T]): T = Await.result(t, Duration.fromSeconds(5)) -@RunWith(classOf[JUnitRunner]) -class AsyncQueueTest extends FunSuite { test("queue pollers") { val q = new AsyncQueue[Int] @@ -132,7 +132,7 @@ class AsyncQueueTest extends FunSuite { q.offer(2) q.offer(3) - assert(1 == Await.result(q.poll())) + assert(1 == await(q.poll())) assert(Return(Queue(2, 3)) == q.drain()) assert(!q.poll().isDefined) @@ -161,4 +161,19 @@ class AsyncQueueTest extends FunSuite { assert(0 == q.size) } + test("size") { + val q = new AsyncQueue[Int]() + assert(q.size == 0) + val f1 = q.poll() + assert(q.size == 0) + q.offer(0) + assert(await(f1) == 0) + assert(q.size == 0) + q.offer(1) + assert(q.size == 1) + q.fail(new Exception(), false) + assert(q.size == 1) + assert(await(q.poll()) == 1) + assert(q.size == 0) + } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/AsyncSemaphoreTest.scala b/util-core/src/test/scala/com/twitter/concurrent/AsyncSemaphoreTest.scala index d6512a0f96..8ee13d57d7 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/AsyncSemaphoreTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/AsyncSemaphoreTest.scala @@ -1,25 +1,29 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util._ -import java.util.concurrent.{ConcurrentLinkedQueue, RejectedExecutionException} -import org.junit.runner.RunWith -import org.scalatest.fixture.FunSpec -import org.scalatest.junit.JUnitRunner +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.{ + ConcurrentLinkedQueue, + CountDownLatch, + Executors, + RejectedExecutionException, + ThreadLocalRandom +} import scala.collection.mutable +import org.scalatest.Outcome +import org.scalatest.funspec.FixtureAnyFunSpec -@RunWith(classOf[JUnitRunner]) -class AsyncSemaphoreTest extends FunSpec { +class AsyncSemaphoreTest extends FixtureAnyFunSpec { class AsyncSemaphoreHelper( val sem: AsyncSemaphore, var count: Int, - val permits: ConcurrentLinkedQueue[Permit] - ) { + val permits: ConcurrentLinkedQueue[Permit]) { def copy( sem: AsyncSemaphore = this.sem, count: Int = this.count, permits: ConcurrentLinkedQueue[Permit] = this.permits - ) = + ): AsyncSemaphoreHelper = new AsyncSemaphoreHelper(sem, count, permits) } @@ -27,7 +31,7 @@ class AsyncSemaphoreTest extends FunSpec { private def await[T](t: Awaitable[T]): T = Await.result(t, 15.seconds) - override def withFixture(test: OneArgTest) = { + override def withFixture(test: OneArgTest): Outcome = { val sem = new AsyncSemaphore(2) val helper = new AsyncSemaphoreHelper(sem, 0, new ConcurrentLinkedQueue[Permit]) withFixture(test.toNoArgTest(helper)) @@ -271,11 +275,11 @@ class AsyncSemaphoreTest extends FunSpec { assert(counter == 2) assert(semHelper.sem.numPermitsAvailable == 0) - queue.dequeue().setValue(Unit) + queue.dequeue().setValue(()) assert(counter == 3) assert(semHelper.sem.numPermitsAvailable == 0) - queue.dequeue().setValue(Unit) + queue.dequeue().setValue(()) assert(semHelper.sem.numPermitsAvailable == 1) queue.dequeue().setException(new RuntimeException("test")) @@ -323,7 +327,7 @@ class AsyncSemaphoreTest extends FunSpec { // sync func is blocked at this point. // But it should be executed as soon as one of the queued up future functions finish - queue.dequeue().setValue(Unit) + queue.dequeue().setValue(()) assert(counter == 3) val result = Await.result(future) assert(result == 3) @@ -362,7 +366,7 @@ class AsyncSemaphoreTest extends FunSpec { // sync func is blocked at this point. // But it should be executed as soon as one of the queued up future functions finish - queue.dequeue().setValue(Unit) + queue.dequeue().setValue(()) assert(counter == 2) assert(Try(Await.result(future)).isThrow) assert(semHelper.sem.numPermitsAvailable == 1) @@ -390,5 +394,34 @@ class AsyncSemaphoreTest extends FunSpec { val msgs = results.collect { case Some(Throw(e)) => e.getMessage } assert(msgs.forall(_ == "woop")) } + + it("behaves properly under contention") { () => + val as = new AsyncSemaphore(5) + val exec = Executors.newCachedThreadPool() + val pool = FuturePool(exec) + val active = new AtomicInteger(0) + try { + val done = + Future.collect { + for (i <- 0 until 10) yield { + pool { + for (j <- 0 until 100) { + val res = + as.acquireAndRunSync { + assert(active.incrementAndGet() <= as.numInitialPermits) + Thread.sleep(ThreadLocalRandom.current().nextInt(10)) + assert(active.decrementAndGet() <= as.numInitialPermits) + i * j + } + assert(await(res) == i * j) + } + } + } + } + await(done) + } finally { + exec.shutdown() + } + } } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/AsyncStreamTest.scala b/util-core/src/test/scala/com/twitter/concurrent/AsyncStreamTest.scala index 649aa4cfed..d24f1fcb7f 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/AsyncStreamTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/AsyncStreamTest.scala @@ -1,770 +1,808 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ -import com.twitter.io.{Buf, Reader} +import com.twitter.conversions.DurationOps._ import com.twitter.util._ -import org.junit.runner.RunWith import org.scalacheck.{Arbitrary, Gen} -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.compat.immutable.LazyList -@RunWith(classOf[JUnitRunner]) -class AsyncStreamTest extends FunSuite with GeneratorDrivenPropertyChecks { +class AsyncStreamTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { import AsyncStream.{mk, of} import AsyncStreamTest._ - test("strict head") { - intercept[Exception] { (undefined: Unit) +:: AsyncStream.empty } - intercept[Exception] { mk(undefined, AsyncStream.empty) } - intercept[Exception] { of(undefined) } + test("++ with a long stream") { + var count = 0 + def genLongStream(len: Int): AsyncStream[Int] = + if (len == 0) { + AsyncStream.of(1) + } else { + count = count + 1 + 1 +:: genLongStream(len - 1) + } + // concat a long stream does not stack overflow + val s = genLongStream(1000000) ++ genLongStream(3) + s.foreach(_ => ()) + val first = count + s.foreach(_ => ()) + // the values are evaluated once + assert(count == first) } - test("lazy tail") { - var forced = false - val s = () +:: { forced = true; AsyncStream.empty[Unit] } - assert(await(s.head) == Some(())) - assert(!forced) - await(s.tail) - assert(forced) - - var forced1 = false - val t = mk((), { forced1 = true; AsyncStream.empty[Unit] }) - assert(await(t.head) == Some(())) - assert(!forced1) - await(t.tail) - assert(forced1) + test("flatMap with a long stream") { + var count = 0 + def genLongStream(len: Int): AsyncStream[Int] = + if (len == 0) { + AsyncStream.of(1) + } else + AsyncStream.of(1).flatMap { i => + count = count + 1 + i +:: genLongStream(len - 1) + } + + // a long stream created via flatMap does not stack overflow + val s = genLongStream(1000000) ++ genLongStream(3) + s.foreach(_ => ()) + val first = count + s.foreach(_ => ()) + // the values are evaluated once + assert(count == first) } - test("call-by-name tail evaluated at most once") { - val p = new Promise[Unit] - val s = () +:: { - if (p.setDone()) of(()) - else AsyncStream.empty[Unit] + // Test all AsyncStream constructors: Empty, FromFuture, Cons, Embed. + seqImpl.foreach { impl => + def fromSeq[A](seq: Seq[A]): AsyncStream[A] = impl.apply(seq) + + test(s"$impl: strict head") { + intercept[Exception] { (undefined: Unit) +:: AsyncStream.empty } + intercept[Exception] { mk(undefined, AsyncStream.empty) } + intercept[Exception] { of(undefined) } } - assert(toSeq(s) == toSeq(s)) - } - test("ops that force tail evaluation") { - def isForced(f: AsyncStream[_] => Future[_]): Unit = { + test(s"$impl: lazy tail") { var forced = false - Await.ready(f(() +:: { forced = true; AsyncStream.empty })) + val s = () +:: { forced = true; AsyncStream.empty[Unit] } + assert(await(s.head) == Some(())) + assert(!forced) + await(s.tail) assert(forced) + + var forced1 = false + val t = mk((), { forced1 = true; AsyncStream.empty[Unit] }) + assert(await(t.head) == Some(())) + assert(!forced1) + await(t.tail) + assert(forced1) } - isForced(_.foldLeft(0)((_, _) => 0)) - isForced(_.foldLeftF(0)((_, _) => Future.value(0))) - isForced(_.tail) - } + test(s"$impl: call-by-name tail evaluated at most once") { + val p = new Promise[Unit] + val s = () +:: { + if (p.setDone()) of(()) + else AsyncStream.empty[Unit] + } + assert(toSeq(s) == toSeq(s)) + } + + test(s"$impl: ops that force tail evaluation") { + def isForced(f: AsyncStream[_] => Future[_]): Unit = { + var forced = false + Await.ready(f(() +:: { forced = true; AsyncStream.empty })) + assert(forced) + } - test("observe: failure") { - val s = 1 +:: 2 +:: (undefined: AsyncStream[Int]) - val (x +: y +: Nil, exc) = await(s.observe()) + isForced(_.foldLeft(0)((_, _) => 0)) + isForced(_.foldLeftF(0)((_, _) => Future.value(0))) + isForced(_.tail) + } - assert(x == 1) - assert(y == 2) - assert(exc.isDefined) - } + test(s"$impl: observe: failure") { + val s = 1 +:: 2 +:: (undefined: AsyncStream[Int]) + val (x +: y +: Nil, exc) = await(s.observe()) - test("observe: no failure") { - val s = 1 +:: 2 +:: AsyncStream.empty[Int] - val (x +: y +: Nil, exc) = await(s.observe()) + assert(x == 1) + assert(y == 2) + assert(exc.isDefined) + } - assert(x == 1) - assert(y == 2) - assert(exc.isEmpty) - } + test(s"$impl: observe: no failure") { + val s = 1 +:: 2 +:: AsyncStream.empty[Int] + val (x +: y +: Nil, exc) = await(s.observe()) - test("fromSeq works on infinite streams") { - def ones: Stream[Int] = 1 #:: ones - assert(toSeq(fromSeq(ones).take(3)) == Seq(1, 1, 1)) - } + assert(x == 1) + assert(y == 2) + assert(exc.isEmpty) + } - test("foreach") { - val x = new Promise[Unit] - val y = new Promise[Unit] + test(s"$impl: fromSeq works on infinite streams") { + def ones: LazyList[Int] = 1 #:: ones + assert(toSeq(fromSeq(ones).take(3)) == Seq(1, 1, 1)) + } - def f() = { x.setDone(); () } - def g() = { y.setDone(); () } + test(s"$impl: foreach") { + val x = new Promise[Unit] + val y = new Promise[Unit] - val s = () +:: f() +:: g() +:: AsyncStream.empty[Unit] - assert(!x.isDefined) - assert(!y.isDefined) + def f() = { x.setDone(); () } + def g() = { y.setDone(); () } - s.foreach(_ => ()) - assert(x.isDefined) - assert(y.isDefined) - } + val s = () +:: f() +:: g() +:: AsyncStream.empty[Unit] + assert(!x.isDefined) + assert(!y.isDefined) - test("lazy ops") { - val p = new Promise[Unit] - val s = () +:: { - p.setDone() - undefined: AsyncStream[Unit] + s.foreach(_ => ()) + assert(x.isDefined) + assert(y.isDefined) } - s.map(x => 0) - assert(!p.isDefined) + test(s"$impl: lazy ops") { + val p = new Promise[Unit] + val s = () +:: { + p.setDone() + undefined: AsyncStream[Unit] + } - s.mapF(x => Future.True) - assert(!p.isDefined) + s.map(x => 0) + assert(!p.isDefined) - s.flatMap(x => of(x)) - assert(!p.isDefined) + s.mapF(x => Future.True) + assert(!p.isDefined) - s.filter(_ => true) - assert(!p.isDefined) + s.flatMap(x => of(x)) + assert(!p.isDefined) - s.withFilter(_ => true) - assert(!p.isDefined) + s.filter(_ => true) + assert(!p.isDefined) - s.take(Int.MaxValue) - assert(!p.isDefined) + s.withFilter(_ => true) + assert(!p.isDefined) - assert(toSeq(s.take(1)) == Seq(())) - assert(!p.isDefined) + s.take(Int.MaxValue) + assert(!p.isDefined) - s.takeWhile(_ => true) - assert(!p.isDefined) + assert(toSeq(s.take(1)) == Seq(())) + assert(!p.isDefined) - s.uncons - assert(!p.isDefined) + s.takeWhile(_ => true) + assert(!p.isDefined) - s.foldRight(Future.Done) { (_, _) => - Future.Done - } - assert(!p.isDefined) + s.uncons + assert(!p.isDefined) - s.scanLeft(Future.Done) { (_, _) => - Future.Done - } - assert(!p.isDefined) + s.foldRight(Future.Done) { (_, _) => Future.Done } + assert(!p.isDefined) - s ++ s - assert(!p.isDefined) + s.scanLeft(Future.Done) { (_, _) => Future.Done } + assert(!p.isDefined) - assert(await(s.head) == Some(())) - assert(!p.isDefined) + s ++ s + assert(!p.isDefined) - intercept[Exception] { await(s.tail).isEmpty } - assert(p.isDefined) - } + assert(await(s.head) == Some(())) + assert(!p.isDefined) - class Ctx[A](ops: AsyncStream[Int] => AsyncStream[A]) { - var once = 0 - val s: AsyncStream[Int] = 2 +:: { - once = once + 1 - if (once > 1) throw new Exception("evaluated more than once") - AsyncStream.of(1) + intercept[Exception] { await(s.tail).isEmpty } + assert(p.isDefined) } - val ss = ops(s) - ss.foreach(_ => ()) - // does not throw - ss.foreach(_ => ()) - } - - test("memoized stream") { - new Ctx(s => s.map(_ => 0)) - new Ctx(s => s.mapF(_ => Future.value(1))) - new Ctx(s => s.flatMap(of(_))) - new Ctx(s => s.filter(_ => true)) - new Ctx(s => s.withFilter(_ => true)) - new Ctx(s => s.take(2)) - new Ctx(s => s.takeWhile(_ => true)) - new Ctx( - s => - s.scanLeft(Future.Done) { (_, _) => - Future.Done - } - ) - new Ctx(s => s ++ s) - } - - // Note: We could use ScalaCheck's Arbitrary[Function1] for some of the tests - // below, however ScalaCheck generates only constant functions which return - // the same value for any input. This makes it quite useless to us. We'll take - // another look since https://github.com/rickynils/scalacheck/issues/136 might - // have solved this issue. + class Ctx[A](ops: AsyncStream[Int] => AsyncStream[A]) { + var once = 0 + val s: AsyncStream[Int] = 2 +:: { + once = once + 1 + if (once > 1) throw new Exception("evaluated more than once") + AsyncStream.of(1) + } - test("map") { - forAll { (s: List[Int]) => - def f(n: Int) = n.toString - assert(toSeq(fromSeq(s).map(f)) == s.map(f)) + val ss = ops(s) + ss.foreach(_ => ()) + // does not throw + ss.foreach(_ => ()) + } + + test(s"$impl: memoized stream") { + new Ctx(s => s.map(_ => 0)) + new Ctx(s => s.mapF(_ => Future.value(1))) + new Ctx(s => s.flatMap(of(_))) + new Ctx(s => s.filter(_ => true)) + new Ctx(s => s.withFilter(_ => true)) + new Ctx(s => s.take(2)) + new Ctx(s => s.takeWhile(_ => true)) + new Ctx(s => s.scanLeft(Future.Done) { (_, _) => Future.Done }) + new Ctx(s => s ++ s) + } + + // Note: We could use ScalaCheck's Arbitrary[Function1] for some of the tests + // below, however ScalaCheck generates only constant functions which return + // the same value for any input. This makes it quite useless to us. We'll take + // another look since https://github.com/rickynils/scalacheck/issues/136 might + // have solved this issue. + + test(s"$impl: map") { + forAll { (s: List[Int]) => + def f(n: Int) = n.toString + assert(toSeq(fromSeq(s).map(f)) == s.map(f)) + } } - } - test("mapF") { - forAll { (s: List[Int]) => - def f(n: Int) = n.toString - val g = f _ andThen Future.value - assert(toSeq(fromSeq(s).mapF(g)) == s.map(f)) + test(s"$impl: mapF") { + forAll { (s: List[Int]) => + def f(n: Int) = n.toString + val g = f _ andThen Future.value + assert(toSeq(fromSeq(s).mapF(g)) == s.map(f)) + } } - } - test("flatMap") { - forAll { (s: List[Int]) => - def f(n: Int) = n.toString - def g(a: Int): AsyncStream[String] = of(f(a)) - def h(a: Int): List[String] = List(f(a)) - assert(toSeq(fromSeq(s).flatMap(g)) == s.flatMap(h)) + test(s"$impl: flatMap") { + forAll { (s: List[Int]) => + def f(n: Int) = n.toString + def g(a: Int): AsyncStream[String] = of(f(a)) + def h(a: Int): List[String] = List(f(a)) + assert(toSeq(fromSeq(s).flatMap(g)) == s.flatMap(h)) + } } - } - test("filter") { - forAll { (s: List[Int]) => - def f(n: Int) = n % 3 == 0 - assert(toSeq(fromSeq(s).filter(f)) == s.filter(f)) + test(s"$impl: filter") { + forAll { (s: List[Int]) => + def f(n: Int) = n % 3 == 0 + assert(toSeq(fromSeq(s).filter(f)) == s.filter(f)) + } } - } - test("++") { - forAll { (a: List[Int], b: List[Int]) => - assert(toSeq(fromSeq(a) ++ fromSeq(b)) == a ++ b) + test(s"$impl: ++") { + forAll { (a: List[Int], b: List[Int]) => assert(toSeq(fromSeq(a) ++ fromSeq(b)) == a ++ b) } } - } - test("++ with a long stream") { - var count = 0 - def genLongStream(len: Int): AsyncStream[Int] = - if (len == 0) { - AsyncStream.of(1) - } else { - count = count + 1 - 1 +:: genLongStream(len - 1) + test(s"$impl: foldRight") { + forAll { (a: List[Int]) => + def f(n: Int, s: String) = (s.toLong + n).toString + def g(q: Int, p: => Future[String]): Future[String] = p.map(f(q, _)) + val m = fromSeq(a).foldRight(Future.value("0"))(g) + assert(await(m) == a.foldRight("0")(f)) } - // concat a long stream does not stack overflow - val s = genLongStream(1000000) ++ genLongStream(3) - s.foreach(_ => ()) - val first = count - s.foreach(_ => ()) - // the values are evaluated once - assert(count == first) - } - - test("foldRight") { - forAll { (a: List[Int]) => - def f(n: Int, s: String) = (s.toLong + n).toString - def g(q: Int, p: => Future[String]): Future[String] = p.map(f(q, _)) - val m = fromSeq(a).foldRight(Future.value("0"))(g) - assert(await(m) == a.foldRight("0")(f)) } - } - test("scanLeft") { - forAll { (a: List[Int]) => - def f(s: String, n: Int) = (s.toLong + n).toString - assert(toSeq(fromSeq(a).scanLeft("0")(f)) == a.scanLeft("0")(f)) + test(s"$impl: scanLeft") { + forAll { (a: List[Int]) => + def f(s: String, n: Int) = (s.toLong + n).toString + assert(toSeq(fromSeq(a).scanLeft("0")(f)) == a.scanLeft("0")(f)) + } } - } - test("scanLeft is eager") { - val never = AsyncStream.fromFuture(Future.never) - val hd = never.scanLeft("hi")((_, _) => ???).head - assert(hd.isDefined) - assert(await(hd) == Some("hi")) - } + test(s"$impl: scanLeft is eager") { + val never = AsyncStream.fromFuture(Future.never) + val hd = never.scanLeft("hi")((_, _) => ???).head + assert(hd.isDefined) + assert(await(hd) == Some("hi")) + } - test("scanLeft is eager for embed") { - val embed = AsyncStream.embed(Future.never) - val hd = embed.scanLeft("hi")((_, _) => ???).head - assert(hd.isDefined) - assert(await(hd) == Some("hi")) - } + test(s"$impl: scanLeft is eager for embed") { + val embed = AsyncStream.embed(Future.never) + val hd = embed.scanLeft("hi")((_, _) => ???).head + assert(hd.isDefined) + assert(await(hd) == Some("hi")) + } - test("foldLeft") { - forAll { (a: List[Int]) => - def f(s: String, n: Int) = (s.toLong + n).toString - assert(await(fromSeq(a).foldLeft("0")(f)) == a.foldLeft("0")(f)) + test(s"$impl: foldLeft") { + forAll { (a: List[Int]) => + def f(s: String, n: Int) = (s.toLong + n).toString + assert(await(fromSeq(a).foldLeft("0")(f)) == a.foldLeft("0")(f)) + } } - } - test("foldLeftF") { - forAll { (a: List[Int]) => - def f(s: String, n: Int) = (s.toLong + n).toString - val g: (String, Int) => Future[String] = (q, p) => Future.value(f(q, p)) - assert(await(fromSeq(a).foldLeftF("0")(g)) == a.foldLeft("0")(f)) + test(s"$impl: foldLeftF") { + forAll { (a: List[Int]) => + def f(s: String, n: Int) = (s.toLong + n).toString + val g: (String, Int) => Future[String] = (q, p) => Future.value(f(q, p)) + assert(await(fromSeq(a).foldLeftF("0")(g)) == a.foldLeft("0")(f)) + } } - } - test("flatten") { - val small = Gen.resize(10, Arbitrary.arbitrary[List[List[Int]]]) - forAll(small) { s => - assert(toSeq(fromSeq(s.map(fromSeq)).flatten) == s.flatten) + test(s"$impl: flatten") { + val small = Gen.resize(10, Arbitrary.arbitrary[List[List[Int]]]) + forAll(small) { s => assert(toSeq(fromSeq(s.map(fromSeq)).flatten) == s.flatten) } } - } - test("head") { - forAll { (a: List[Int]) => - assert(await(fromSeq(a).head) == a.headOption) + test(s"$impl: head") { + forAll { (a: List[Int]) => assert(await(fromSeq(a).head) == a.headOption) } } - } - test("isEmpty") { - val s = AsyncStream.of(1) - val tail = await(s.tail) - assert(tail == None) - } + test(s"$impl: isEmpty") { + val s = AsyncStream.of(1) + val tail = await(s.tail) + assert(tail == None) + } - test("tail") { - forAll(Gen.nonEmptyListOf(Arbitrary.arbitrary[Int])) { (a: List[Int]) => - val tail = await(fromSeq(a).tail) - a.tail match { - case Nil => assert(tail == None) - case _ => assert(toSeq(tail.get) == a.tail) + test(s"$impl: tail") { + forAll(Gen.nonEmptyListOf(Arbitrary.arbitrary[Int])) { (a: List[Int]) => + val tail = await(fromSeq(a).tail) + a.tail match { + case Nil => assert(tail == None) + case _ => assert(toSeq(tail.get) == a.tail) + } } } - } - test("uncons") { - assert(await(AsyncStream.empty.uncons) == None) - forAll(Gen.nonEmptyListOf(Arbitrary.arbitrary[Int])) { (a: List[Int]) => - val Some((h, t)) = await(fromSeq(a).uncons) - assert(h == a.head) - assert(toSeq(t()) == a.tail) + test(s"$impl: uncons") { + assert(await(AsyncStream.empty.uncons) == None) + forAll(Gen.nonEmptyListOf(Arbitrary.arbitrary[Int])) { (a: List[Int]) => + val Some((h, t)) = await(fromSeq(a).uncons) + assert(h == a.head) + assert(toSeq(t()) == a.tail) + } } - } - test("take") { - forAll(genListAndN) { - case (as, n) => - assert(toSeq(fromSeq(as).take(n)) == as.take(n)) + test(s"$impl: take") { + forAll(genListAndN) { + case (as, n) => + assert(toSeq(fromSeq(as).take(n)) == as.take(n)) + } } - } - test("drop") { - forAll(genListAndN) { - case (as, n) => - assert(toSeq(fromSeq(as).drop(n)) == as.drop(n)) + test(s"$impl: drop") { + forAll(genListAndN) { + case (as, n) => + assert(toSeq(fromSeq(as).drop(n)) == as.drop(n)) + } } - } - test("takeWhile") { - forAll(genListAndSentinel) { - case (as, x) => - assert(toSeq(fromSeq(as).takeWhile(_ != x)) == as.takeWhile(_ != x)) + test(s"$impl: takeWhile") { + forAll(genListAndSentinel) { + case (as, x) => + assert(toSeq(fromSeq(as).takeWhile(_ != x)) == as.takeWhile(_ != x)) + } } - } - test("dropWhile") { - forAll(genListAndSentinel) { - case (as, x) => - assert(toSeq(fromSeq(as).dropWhile(_ != x)) == as.dropWhile(_ != x)) + test(s"$impl: dropWhile") { + forAll(genListAndSentinel) { + case (as, x) => + assert(toSeq(fromSeq(as).dropWhile(_ != x)) == as.dropWhile(_ != x)) + } } - } - test("toSeq") { - forAll { (as: List[Int]) => - assert(await(fromSeq(as).toSeq()) == as) + test(s"$impl: toSeq") { + forAll { (as: List[Int]) => assert(await(fromSeq(as).toSeq()) == as) } } - } - test("identity") { - val small = Gen.resize(10, Arbitrary.arbitrary[List[Int]]) - forAll(small) { s => - val a = fromSeq(s) - def f(x: Int) = x +:: a + test(s"$impl: identity") { + val small = Gen.resize(10, Arbitrary.arbitrary[List[Int]]) + forAll(small) { s => + val a = fromSeq(s) + def f(x: Int) = x +:: a - assert(toSeq(of(1).flatMap(f)) == toSeq(f(1))) - assert(toSeq(a.flatMap(of)) == toSeq(a)) + assert(toSeq(of(1).flatMap(f)) == toSeq(f(1))) + assert(toSeq(a.flatMap(of)) == toSeq(a)) + } } - } - test("associativity") { - val small = Gen.resize(10, Arbitrary.arbitrary[List[Int]]) - forAll(small, small, small) { (s, t, u) => - val a = fromSeq(s) - val b = fromSeq(t) - val c = fromSeq(u) + test(s"$impl: associativity") { + val small = Gen.resize(10, Arbitrary.arbitrary[List[Int]]) + forAll(small, small, small) { (s, t, u) => + val a = fromSeq(s) + val b = fromSeq(t) + val c = fromSeq(u) - def f(x: Int) = x +:: b - def g(x: Int) = x +:: c + def f(x: Int) = x +:: b + def g(x: Int) = x +:: c - val v = a.flatMap(f).flatMap(g) - val w = a.flatMap(x => f(x).flatMap(g)) - assert(toSeq(v) == toSeq(w)) + val v = a.flatMap(f).flatMap(g) + val w = a.flatMap(x => f(x).flatMap(g)) + assert(toSeq(v) == toSeq(w)) + } } - } - test("buffer() works like Seq.splitAt") { - forAll { (items: List[Char], bufferSize: Int) => - val (expectedBuffer, expectedRest) = items.splitAt(bufferSize) - val (buffer, rest) = await(fromSeq(items).buffer(bufferSize)) - assert(expectedBuffer == buffer) - assert(expectedRest == toSeq(rest())) + test(s"$impl: buffer() works like Seq.splitAt") { + forAll { (items: List[Char], bufferSize: Int) => + val (expectedBuffer, expectedRest) = items.splitAt(bufferSize) + val (buffer, rest) = await(fromSeq(items).buffer(bufferSize)) + assert(expectedBuffer == buffer) + assert(expectedRest == toSeq(rest())) + } } - } - test("buffer() has the same properties as take() and drop()") { - // We need items to be non-empty, because AsyncStream.empty ++ - // forces the future to be created. - val gen = Gen.zip(Gen.nonEmptyListOf(Arbitrary.arbitrary[Char]), Arbitrary.arbitrary[Int]) - - forAll(gen) { - case (items, n) => - var forced1 = false - val stream1 = fromSeq(items) ++ { forced1 = true; AsyncStream.empty[Char] } - var forced2 = false - val stream2 = fromSeq(items) ++ { forced2 = true; AsyncStream.empty[Char] } - - val takeResult = toSeq(stream2.take(n)) - val (bufferResult, bufferRest) = await(stream1.buffer(n)) - assert(takeResult == bufferResult) - - // Strictness property: we should only need to force the full - // stream if we asked for more items that were present in the - // stream. - assert(forced1 == (n > items.size)) - assert(forced1 == forced2) - val wasForced = forced1 - - // Strictness property: Since AsyncStream contains a Future - // rather than a thunk, we need to evaluate the next element in - // order to get the result of drop and the rest of the stream - // after buffering. - val bufferTail = bufferRest() - val dropTail = stream2.drop(n) - assert(forced1 == (n >= items.size)) - assert(forced1 == forced2) - - // This is the only case that should have caused the item to be forced. - assert((wasForced == forced1) || n == items.size) - - // Forcing the rest of the sequence should always cause evaluation. - assert(toSeq(bufferTail) == toSeq(dropTail)) - assert(forced1) - assert(forced2) + test(s"$impl: buffer() has the same properties as take() and drop()") { + // We need items to be non-empty, because AsyncStream.empty ++ + // forces the future to be created. + val gen = Gen.zip(Gen.nonEmptyListOf(Arbitrary.arbitrary[Char]), Arbitrary.arbitrary[Int]) + + forAll(gen) { + case (items, n) => + var forced1 = false + val stream1 = fromSeq(items) ++ { forced1 = true; AsyncStream.empty[Char] } + var forced2 = false + val stream2 = fromSeq(items) ++ { forced2 = true; AsyncStream.empty[Char] } + + val takeResult = toSeq(stream2.take(n)) + val (bufferResult, bufferRest) = await(stream1.buffer(n)) + assert(takeResult == bufferResult) + + // Strictness property: we should only need to force the full + // stream if we asked for more items that were present in the + // stream. + assert(forced1 == (n > items.size)) + assert(forced1 == forced2) + val wasForced = forced1 + + // Strictness property: Since AsyncStream contains a Future + // rather than a thunk, we need to evaluate the next element in + // order to get the result of drop and the rest of the stream + // after buffering. + val bufferTail = bufferRest() + val dropTail = stream2.drop(n) + assert(forced1 == (n >= items.size)) + assert(forced1 == forced2) + + // This is the only case that should have caused the item to be forced. + assert((wasForced == forced1) || n == items.size) + + // Forcing the rest of the sequence should always cause evaluation. + assert(toSeq(bufferTail) == toSeq(dropTail)) + assert(forced1) + assert(forced2) + } } - } - test("grouped() works like Seq.grouped") { - forAll { (items: Seq[Char], groupSize: Int) => - // This is a Try so that we can test that bad inputs act the - // same. (Zero or negative group sizes throw the same - // exception.) - val expected = Try(items.grouped(groupSize).toSeq) - val actual = Try(toSeq(fromSeq(items).grouped(groupSize))) - - // If they are both exceptions, then pass if the exceptions are - // the same type (don't require them to define equality or have - // the same exception message) - (actual, expected) match { - case (Throw(e1), Throw(e2)) => assert(e1.getClass == e2.getClass) - case _ => assert(actual == expected) + test(s"$impl: grouped() works like Seq.grouped") { + forAll { (items: Seq[Char], groupSize: Int) => + // This is a Try so that we can test that bad inputs act the + // same. (Zero or negative group sizes throw the same + // exception.) + val expected = Try(items.grouped(groupSize).toSeq) + val actual = Try(toSeq(fromSeq(items).grouped(groupSize))) + + // If they are both exceptions, then pass if the exceptions are + // the same type (don't require them to define equality or have + // the same exception message) + (actual, expected) match { + case (Throw(e1), Throw(e2)) => assert(e1.getClass == e2.getClass) + case _ => assert(actual == expected) + } } } - } - test("grouped should be lazy") { - val gen = - for { - // We need items to be non-empty, because AsyncStream.empty ++ - // forces the future to be created. - items <- Gen.nonEmptyListOf(Arbitrary.arbitrary[Char]) - - // We need to make sure that the chunk size (1) is valid and (2) - // is short enough that forcing the first group does not force - // the exception. - groupSize <- Gen.chooseNum(1, items.size) - } yield (items, groupSize) - - forAll(gen) { - case (items, groupSize) => - var forced = false - val stream: AsyncStream[Char] = fromSeq(items) ++ { forced = true; AsyncStream.empty } - - val expected = items.grouped(groupSize).toSeq.headOption - // This will take up to items.size items from the stream. This - // does not require forcing the tail. - val actual = await(stream.grouped(groupSize).head) - assert(actual == expected) - assert(!forced) - val expectedChunks = items.grouped(groupSize).toSeq - val allChunks = toSeq(stream.grouped(groupSize)) - assert(allChunks == expectedChunks) - assert(forced) + test(s"$impl: grouped should be lazy") { + val gen = + for { + // We need items to be non-empty, because AsyncStream.empty ++ + // forces the future to be created. + items <- Gen.nonEmptyListOf(Arbitrary.arbitrary[Char]) + + // We need to make sure that the chunk size (1) is valid and (2) + // is short enough that forcing the first group does not force + // the exception. + groupSize <- Gen.chooseNum(1, items.size) + } yield (items, groupSize) + + forAll(gen) { + case (items, groupSize) => + var forced = false + val stream: AsyncStream[Char] = fromSeq(items) ++ { forced = true; AsyncStream.empty } + + val expected = items.grouped(groupSize).toSeq.headOption + // This will take up to items.size items from the stream. This + // does not require forcing the tail. + val actual = await(stream.grouped(groupSize).head) + assert(actual == expected) + assert(!forced) + val expectedChunks = items.grouped(groupSize).toSeq + val allChunks = toSeq(stream.grouped(groupSize)) + assert(allChunks == expectedChunks) + assert(forced) + } } - } - test("mapConcurrent preserves items") { - forAll(Arbitrary.arbitrary[List[Int]], Gen.choose(1, 10)) { (xs, conc) => - assert(toSeq(AsyncStream.fromSeq(xs).mapConcurrent(conc)(Future.value)).sorted == xs.sorted) + test(s"$impl: mapConcurrent preserves items") { + forAll(Arbitrary.arbitrary[List[Int]], Gen.choose(1, 10)) { (xs, conc) => + assert(toSeq(fromSeq(xs).mapConcurrent(conc)(Future.value)).sorted == xs.sorted) + } } - } - test("mapConcurrent makes progress when an item is blocking") { - forAll(Arbitrary.arbitrary[List[Int]], Gen.choose(2, 10)) { (xs, conc) => - // This promise is not satisfied, which would block the evaluation - // of .map, and should not block .mapConcurrent when conc > 1 - val first = new Promise[Int] - - // This function will return a blocking future the first time it - // is called and an immediately-available future thereafter. - var used = false - def f(x: Int) = - if (used) { - Future.value(x) - } else { - used = true - first - } + test(s"$impl: mapConcurrent makes progress when an item is blocking") { + forAll(Arbitrary.arbitrary[List[Int]], Gen.choose(2, 10)) { (xs, conc) => + // This promise is not satisfied, which would block the evaluation + // of .map, and should not block .mapConcurrent when conc > 1 + val first = new Promise[Int] + + // This function will return a blocking future the first time it + // is called and an immediately-available future thereafter. + var used = false + def f(x: Int) = + if (used) { + Future.value(x) + } else { + used = true + first + } - // Concurrently map over the stream. The whole stream should be - // available, except for one item which is still blocked. - val mapped = AsyncStream.fromSeq(xs).mapConcurrent(conc)(f) + // Concurrently map over the stream. The whole stream should be + // available, except for one item which is still blocked. + val mapped = fromSeq(xs).mapConcurrent(conc)(f) - // All but the first value, which is still blocking, has been returned - assert(toSeq(mapped.take(xs.length - 1)).sorted == xs.drop(1).sorted) + // All but the first value, which is still blocking, has been returned + assert(toSeq(mapped.take(xs.length - 1)).sorted == xs.drop(1).sorted) - if (xs.nonEmpty) { - // The stream as a whole is still blocking on the unsatisfied promise - assert(!mapped.foreach(_ => ()).isDefined) + if (xs.nonEmpty) { + // The stream as a whole is still blocking on the unsatisfied promise + assert(!mapped.foreach(_ => ()).isDefined) - // Unblock the first value - first.setValue(xs.head) - } + // Unblock the first value + first.setValue(xs.head) + } - // Now the whole stream should be available and should contain all - // of the items, ignoring order (but preserving repetition) - assert(mapped.foreach(_ => ()).isDefined) - assert(toSeq(mapped).sorted == xs.sorted) + // Now the whole stream should be available and should contain all + // of the items, ignoring order (but preserving repetition) + assert(mapped.foreach(_ => ()).isDefined) + assert(toSeq(mapped).sorted == xs.sorted) + } } - } - test("mapConcurrent is lazy once it reaches its concurrency limit") { - forAll(Gen.choose(2, 10), Arbitrary.arbitrary[Seq[Int]]) { (conc, xs) => - val q = new scala.collection.mutable.Queue[Promise[Unit]] + test(s"$impl: mapConcurrent is lazy once it reaches its concurrency limit") { + forAll(Gen.choose(2, 10), Arbitrary.arbitrary[Seq[Int]]) { (conc, xs) => + val q = new scala.collection.mutable.Queue[Promise[Unit]] - val mapped = - AsyncStream.fromSeq(xs).mapConcurrent(conc) { _ => - val p = new Promise[Unit] - q.enqueue(p) - p - } + val mapped = + fromSeq(xs).mapConcurrent(conc) { _ => + val p = new Promise[Unit] + q.enqueue(p) + p + } - // If there are at least `conc` items in the queue, then we should - // have started exactly `conc` of them. Otherwise, we should have - // started all of them. - assert(q.size == conc.min(xs.size)) + // If there are at least `conc` items in the queue, then we should + // have started exactly `conc` of them. Otherwise, we should have + // started all of them. + assert(q.size == conc.min(xs.size)) - if (xs.nonEmpty) { - assert(!mapped.head.isDefined) + if (xs.nonEmpty) { + assert(!mapped.head.isDefined) - val p = q.dequeue() - p.setDone() - } + val p = q.dequeue() + p.setDone() + } - // Satisfying that promise makes the head of the queue available. - assert(mapped.head.isDefined) + // Satisfying that promise makes the head of the queue available. + assert(mapped.head.isDefined) - if (xs.size > 1) { - // We do not add another element to the queue until the next - // element is forced. - assert(q.size == (conc.min(xs.size) - 1)) + if (xs.size > 1) { + // We do not add another element to the queue until the next + // element is forced. + assert(q.size == (conc.min(xs.size) - 1)) - val tl = mapped.drop(1) - assert(!tl.head.isDefined) + val tl = mapped.drop(1) + assert(!tl.head.isDefined) - // Forcing the next element of the queue causes us to enqueue - // one more element (if there are more elements to enqueue) - assert(q.size == conc.min(xs.size - 1)) + // Forcing the next element of the queue causes us to enqueue + // one more element (if there are more elements to enqueue) + assert(q.size == conc.min(xs.size - 1)) - val p = q.dequeue() - p.setDone() + val p = q.dequeue() + p.setDone() - // Satisfying that promise causes the head to be available. - assert(tl.head.isDefined) + // Satisfying that promise causes the head to be available. + assert(tl.head.isDefined) + } } } - } - - test("mapConcurrent makes progress, even with blocking streams and blocking work") { - val gen = - Gen.zip( - Gen.choose(0, 10).label("numActions"), - Gen.choose(0, 10).flatMap(Gen.listOfN(_, Arbitrary.arbitrary[Int])), - Gen.choose(1, 11).label("concurrency") - ) - forAll(gen) { - case (numActions, items, concurrency) => - val input: AsyncStream[Int] = - AsyncStream.fromSeq(items) ++ AsyncStream.fromFuture(Future.never) - - var workStarted = 0 - var workFinished = 0 - val result = - input.mapConcurrent(concurrency) { i => - workStarted += 1 - if (workFinished < numActions) { - workFinished += 1 - Future.value(i) - } else { - // After numActions evaluations, return a Future that - // will never be satisfied. - Future.never + test(s"$impl: mapConcurrent makes progress, even with blocking streams and blocking work") { + val gen = + Gen.zip( + Gen.choose(0, 10).label("numActions"), + Gen.choose(0, 10).flatMap(Gen.listOfN(_, Arbitrary.arbitrary[Int])), + Gen.choose(1, 11).label("concurrency") + ) + + forAll(gen) { + case (numActions, items, concurrency) => + val input: AsyncStream[Int] = + fromSeq(items) ++ AsyncStream.fromFuture(Future.never) + + var workStarted = 0 + var workFinished = 0 + val result = + input.mapConcurrent(concurrency) { i => + workStarted += 1 + if (workFinished < numActions) { + workFinished += 1 + Future.value(i) + } else { + // After numActions evaluations, return a Future that + // will never be satisfied. + Future.never + } } - } - // How much work should have been started by mapConcurrent. - val expectedStarted = items.size.min(concurrency) - assert(workStarted == expectedStarted, "work started") + // How much work should have been started by mapConcurrent. + val expectedStarted = items.size.min(concurrency) + assert(workStarted == expectedStarted, "work started") - val expectedFinished = numActions.min(expectedStarted) - assert(workFinished == expectedFinished, "expected finished") + val expectedFinished = numActions.min(expectedStarted) + assert(workFinished == expectedFinished, "expected finished") - // Make sure that all of the finished items are now - // available. (As a side-effect, this will force more work to - // be done if concurrency was the limiting factor.) - val completed = toSeq(result.take(workFinished)).sorted - val expectedCompleted = items.take(expectedFinished).sorted - assert(completed == expectedCompleted) + // Make sure that all of the finished items are now + // available. (As a side-effect, this will force more work to + // be done if concurrency was the limiting factor.) + val completed = toSeq(result.take(workFinished)) + assert(completed.toSet.subsetOf(items.toSet)) + } } - } - - test("fromReader") { - forAll { l: List[Byte] => - val buf = Buf.ByteArray.Owned(l.toArray) - val as = AsyncStream.fromReader(Reader.fromBuf(buf), chunkSize = 1) - assert(toSeq(as).map(b => Buf.ByteArray.Owned.extract(b).head) == l) + test(s"$impl: mapConcurrent tail laziness") { + val failAfter1 = 1 +:: undefined[AsyncStream[Int]] + val justOne = failAfter1.mapConcurrent(1)(Future.value).take(1) + assert(await(justOne.toSeq()) == Seq(1)) } - } - test("sum") { - forAll { xs: List[Int] => - assert(xs.sum == await(AsyncStream.fromSeq(xs).sum)) + test(s"$impl: sum") { + forAll { (xs: List[Int]) => assert(xs.sum == await(fromSeq(xs).sum)) } } - } - test("size") { - forAll { xs: List[Int] => - assert(xs.size == await(AsyncStream.fromSeq(xs).size)) + test(s"$impl: size") { + forAll { (xs: List[Int]) => assert(xs.size == await(fromSeq(xs).size)) } } - } - test("force") { - forAll { xs: List[Int] => - val p = new Promise[Unit] - // The promise will be defined iff the tail is forced. - val s = AsyncStream.fromSeq(xs) ++ { p.setDone(); AsyncStream.empty } + test(s"$impl: force") { + forAll { (xs: List[Int]) => + val p = new Promise[Unit] + // The promise will be defined iff the tail is forced. + val s = fromSeq(xs) ++ { p.setDone(); AsyncStream.empty } - // If the input is empty, then the tail will be forced right away. - assert(p.isDefined == xs.isEmpty) + // If the input is empty, then the tail will be forced right away. + assert(p.isDefined == xs.isEmpty) - // Unconditionally force the whole stream - await(s.force) - assert(p.isDefined) + // Unconditionally force the whole stream + await(s.force) + assert(p.isDefined) + } } - } - test("withEffect") { - forAll(genListAndN) { - case (xs, n) => - var i = 0 - val s = AsyncStream.fromSeq(xs).withEffect(_ => i += 1) + test(s"$impl: withEffect") { + forAll(genListAndN) { + case (xs, n) => + var i = 0 + val s = fromSeq(xs).withEffect(_ => i += 1) - // Is lazy on initial application (with the exception of the first element) - assert(i == (if (xs.isEmpty) 0 else 1)) + // Is lazy on initial application (with the exception of the first element) + assert(i == (if (xs.isEmpty) 0 else 1)) - // Is lazy when consuming the stream - await(s.take(n).force) + // Is lazy when consuming the stream + await(s.take(n).force) - // If the list is empty, no effects should occur. If the list is - // non-empty, the effect will occur for the first item right away, - // since the head is not lazy. Otherwise, we expect the same - // number of effects as items demanded. - val expected = if (xs.isEmpty) 0 else 1.max(xs.length.min(n)) - assert(i == expected) + // If the list is empty, no effects should occur. If the list is + // non-empty, the effect will occur for the first item right away, + // since the head is not lazy. Otherwise, we expect the same + // number of effects as items demanded. + val expected = if (xs.isEmpty) 0 else 1.max(xs.length.min(n)) + assert(i == expected) - // Preserves the elements in the stream - assert(toSeq(s) == xs) + // Preserves the elements in the stream + assert(toSeq(s) == xs) + } } - } - test("merge generates a stream equal to all input streams") { - forAll { (lists: Seq[List[Int]]) => - val streams = lists.map(fromSeq) - val merged = AsyncStream.merge(streams: _*) + test(s"$impl: merge generates a stream equal to all input streams") { + forAll { (lists: Seq[List[Int]]) => + val streams = lists.map(fromSeq) + val merged = AsyncStream.merge(streams: _*) - val input = AsyncStream(streams: _*).flatten + val input = AsyncStream(streams: _*).flatten - assert(toSeq(input).sorted == toSeq(merged).sorted) + assert(toSeq(input).sorted == toSeq(merged).sorted) + } } - } - test("merge fails the result stream if an input stream fails") { - forAll() { (lists: Seq[List[Int]]) => - val s = mk(1, undefined: AsyncStream[Int]) - val streams = s +: lists.map(fromSeq) - val merged = AsyncStream.merge(streams: _*) + test(s"$impl: merge fails the result stream if an input stream fails") { + forAll() { (lists: Seq[List[Int]]) => + val s = mk(1, undefined: AsyncStream[Int]) + val streams = s +: lists.map(fromSeq) + val merged = AsyncStream.merge(streams: _*) - intercept[Exception](toSeq(merged)) + intercept[Exception](toSeq(merged)) + } } - } - test("merged stream contains elements as they become available from input streams") { - forAll { (input: List[Int]) => - val promises = List.fill(input.size)(Promise[Int]()) + test(s"$impl: merged stream contains elements as they become available from input streams") { + forAll { (input: List[Int]) => + val promises = List.fill(input.size)(Promise[Int]()) - // grouped into lists of 10 elements each - val grouped = promises.grouped(10).toList + // grouped into lists of 10 elements each + val grouped = promises.grouped(10).toList - // merged list of streams - val streams = grouped.map(fromSeq(_).flatMap(AsyncStream.fromFuture)) - val merged = AsyncStream.merge(streams: _*).toSeq() + // merged list of streams + val streams = grouped.map(fromSeq(_).flatMap(AsyncStream.fromFuture)) + val merged = AsyncStream.merge(streams: _*).toSeq() - // build an interleaved list of the promises for the stream - // [s1(1), s2(1), s3(1), s1(2), s2(2), s3(2), ...] - val interleavedHeads = grouped.flatMap(_.zipWithIndex).sortBy(_._2).map(_._1) - interleavedHeads.zip(input).foreach { - case (p, i) => - p.update(Return(i)) + // build an interleaved list of the promises for the stream + // [s1(1), s2(1), s3(1), s1(2), s2(2), s3(2), ...] + val interleavedHeads = grouped.flatMap(_.zipWithIndex).sortBy(_._2).map(_._1) + interleavedHeads.zip(input).foreach { + case (p, i) => + p.update(Return(i)) + } + + assert(Await.result(merged) == input) } + } - assert(Await.result(merged) == input) + test(s"$impl: exception produces a failed stream") { + intercept[Exception]( + toSeq(AsyncStream.exception(new Exception())) + ) } - } - test("exception produces a failed stream") { - intercept[Exception]( - toSeq(AsyncStream.exception(new Exception())) - ) - } + test(s"$impl: exception eventually produces a failed stream") { + forAll { (xs: List[Int]) => + val stream = fromSeq(xs) ++ AsyncStream.exception(new Exception()) + intercept[Exception](toSeq(stream)) + } + } + + test(s"$impl: flatMap is stack-safe") { + val n = 10000 - test("exception eventually produces a failed stream") { - forAll { (xs: List[Int]) => - val stream = fromSeq(xs) ++ AsyncStream.exception(new Exception()) - intercept[Exception](toSeq(stream)) + def tailRecM[A, B](a: A)(f: A => AsyncStream[Either[A, B]]): AsyncStream[B] = + f(a).flatMap { + case Right(b) => AsyncStream.of(b) + case Left(a) => tailRecM(a)(f) + } + + val stream = tailRecM(0) { i => AsyncStream.of(if (i < n) Left(i + 1) else Right(i)) } + + assert(Await.result(stream.toSeq()) == Seq(n)) } } } private object AsyncStreamTest { - val genListAndN = for { + val genListAndN: Gen[Tuple2[List[Int], Int]] = for { as <- Arbitrary.arbitrary[List[Int]] n <- Gen.choose(0, as.length) } yield (as, n) - val genListAndSentinel = for { + val genListAndSentinel: Gen[Tuple2[List[Int], Int]] = for { as <- Arbitrary.arbitrary[List[Int]] if as.nonEmpty n <- Gen.choose(0, as.length - 1) } yield (as, as(n)) - def await[T](fut: Future[T]) = Await.result(fut, 100.milliseconds) + def await[T](fut: Future[T]): T = Await.result(fut, 100.milliseconds) def undefined[A]: A = throw new Exception def toSeq[A](s: AsyncStream[A]): Seq[A] = await(s.toSeq()) - def fromSeq[A](s: Seq[A]): AsyncStream[A] = - // Test all AsyncStream constructors: Empty, FromFuture, Cons, Embed. - s match { + sealed trait FromSeq { + def apply[A](seq: Seq[A]): AsyncStream[A] + } + + private case object Cons extends FromSeq { + def apply[A](seq: Seq[A]): AsyncStream[A] = seq match { + case Nil => AsyncStream.empty + case a +: as => a +:: apply(as) + } + } + + private case object EmbeddedCons extends FromSeq { + def apply[A](seq: Seq[A]): AsyncStream[A] = seq match { + case Nil => AsyncStream.embed(Future.value(AsyncStream.empty)) + case a +: as => AsyncStream.embed(Future.value(a +:: apply(as))) + } + } + + private case object OfFuture extends FromSeq { + def apply[A](seq: Seq[A]): AsyncStream[A] = seq match { case Nil => AsyncStream.empty case a +: Nil => AsyncStream.of(a) - case a +: b +: Nil => AsyncStream.embed(Future.value(a +:: AsyncStream.of(b))) - case a +: as => a +:: fromSeq(as) + case a +: as => a +:: apply(as) } -} + } + + private case object EmbeddedFuture extends FromSeq { + def apply[A](seq: Seq[A]): AsyncStream[A] = seq match { + case Nil => AsyncStream.embed(Future.value(AsyncStream.empty)) + case a +: Nil => AsyncStream.embed(Future.value(AsyncStream.of(a))) + case a +: as => AsyncStream.embed(Future.value(a +:: apply(as))) + } + } + + def seqImpl: Seq[FromSeq] = Seq(Cons, EmbeddedCons, OfFuture, EmbeddedFuture) +} \ No newline at end of file diff --git a/util-core/src/test/scala/com/twitter/concurrent/BUILD b/util-core/src/test/scala/com/twitter/concurrent/BUILD new file mode 100644 index 0000000000..06206406c7 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/concurrent/BUILD @@ -0,0 +1,23 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/test/java/com/twitter/util:object_size_calculator", + ], +) diff --git a/util-core/src/test/scala/com/twitter/concurrent/BrokerTest.scala b/util-core/src/test/scala/com/twitter/concurrent/BrokerTest.scala index 3a91785c52..00a7c674fd 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/BrokerTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/BrokerTest.scala @@ -1,119 +1,114 @@ package com.twitter.concurrent -import com.twitter.util.{Await, Return, ObjectSizeCalculator} -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class BrokerTest extends WordSpec { - "Broker" should { - "send data (send, recv)" in { - val br = new Broker[Int] - val sendF = br.send(123).sync() - assert(sendF.isDefined == false) - val recvF = br.recv.sync() - assert(recvF.isDefined == true) - assert(Await.result(recvF) == 123) - assert(sendF.isDefined == true) - } +import com.twitter.util.Await +import com.twitter.util.ObjectSizeCalculator +import com.twitter.util.Return +import org.scalatest.funsuite.AnyFunSuite + +class BrokerTest extends AnyFunSuite { + test("Broker should send data (send, recv)") { + val br = new Broker[Int] + val sendF = br.send(123).sync() + assert(sendF.isDefined == false) + val recvF = br.recv.sync() + assert(recvF.isDefined == true) + assert(Await.result(recvF) == 123) + assert(sendF.isDefined == true) + } - "send data (recv, send)" in { - val br = new Broker[Int] - val recvF = br.recv.sync() - assert(recvF.isDefined == false) - val sendF = br.send(123).sync() - assert(sendF.isDefined == true) - assert(recvF.isDefined == true) + test("Broker should send data (recv, send)") { + val br = new Broker[Int] + val recvF = br.recv.sync() + assert(recvF.isDefined == false) + val sendF = br.send(123).sync() + assert(sendF.isDefined == true) + assert(recvF.isDefined == true) - assert(Await.result(recvF) == 123) - } + assert(Await.result(recvF) == 123) + } - "queue receivers (recv, recv, send, send)" in { - val br = new Broker[Int] - val r0, r1 = br.recv.sync() - assert(r0.isDefined == false) - assert(r1.isDefined == false) - val s = br.send(123) - assert(s.sync().poll == Some(Return.Unit)) - assert(r0.poll == Some(Return(123))) - assert(r1.isDefined == false) - assert(s.sync().poll == Some(Return.Unit)) - assert(r1.poll == Some(Return(123))) - assert(s.sync().isDefined == false) - } + test("Broker should queue receivers (recv, recv, send, send)") { + val br = new Broker[Int] + val r0, r1 = br.recv.sync() + assert(r0.isDefined == false) + assert(r1.isDefined == false) + val s = br.send(123) + assert(s.sync().poll == Some(Return.Unit)) + assert(r0.poll == Some(Return(123))) + assert(r1.isDefined == false) + assert(s.sync().poll == Some(Return.Unit)) + assert(r1.poll == Some(Return(123))) + assert(s.sync().isDefined == false) + } - "queue senders (send, send, recv, recv)" in { - val br = new Broker[Int] - val s0, s1 = br.send(123).sync() - assert(s0.isDefined == false) - assert(s1.isDefined == false) - val r = br.recv - assert(r.sync().poll == Some(Return(123))) - assert(s0.poll == Some(Return.Unit)) - assert(s1.isDefined == false) - assert(r.sync().poll == Some(Return(123))) - assert(s1.poll == Some(Return.Unit)) - assert(r.sync().isDefined == false) - } + test("Broker should queue senders (send, send, recv, recv)") { + val br = new Broker[Int] + val s0, s1 = br.send(123).sync() + assert(s0.isDefined == false) + assert(s1.isDefined == false) + val r = br.recv + assert(r.sync().poll == Some(Return(123))) + assert(s0.poll == Some(Return.Unit)) + assert(s1.isDefined == false) + assert(r.sync().poll == Some(Return(123))) + assert(s1.poll == Some(Return.Unit)) + assert(r.sync().isDefined == false) + } - "interrupts" should { - "removes queued receiver" in { - val br = new Broker[Int] - val recvF = br.recv.sync() - recvF.raise(new Exception) - assert(br.send(123).sync().poll == None) - assert(recvF.poll == None) - } - - "removes queued sender" in { - val br = new Broker[Int] - val sendF = br.send(123).sync() - sendF.raise(new Exception) - assert(br.recv.sync().poll == None) - assert(sendF.poll == None) - } - - "doesn't result in space leaks" in { - val br = new Broker[Int] - - assert(Offer.select(Offer.const(1), br.recv).poll == Some(Return(1))) - val initial = ObjectSizeCalculator.getObjectSize(br) - - for (_ <- 0 until 1000) { - assert(Offer.select(Offer.const(1), br.recv).poll == Some(Return(1))) - assert(ObjectSizeCalculator.getObjectSize(br) == initial) - } - } - - "works with orElse" in { - val b0, b1 = new Broker[Int] - - val o = b0.recv orElse b1.recv - val f = o.sync() - assert(f.isDefined == false) - - val sendf0 = b0.send(12).sync() - assert(sendf0.isDefined == false) - val sendf1 = b1.send(32).sync() - assert(sendf1.isDefined == true) - assert(f.poll == Some(Return(32))) - - assert(o.sync().poll == Some(Return(12))) - assert(sendf0.poll == Some(Return.Unit)) - } - } + test("Broker interrupts should remove queued receiver") { + val br = new Broker[Int] + val recvF = br.recv.sync() + recvF.raise(new Exception) + assert(br.send(123).sync().poll == None) + assert(recvF.poll == None) + } - "integrate" in { - val br = new Broker[Int] - val offer = Offer.choose(br.recv, Offer.const(999)) - assert(offer.sync().poll == Some(Return(999))) + test("Broker interrupts should removes queued sender") { + val br = new Broker[Int] + val sendF = br.send(123).sync() + sendF.raise(new Exception) + assert(br.recv.sync().poll == None) + assert(sendF.poll == None) + } - val item = br.recv.sync() - assert(item.isDefined == false) + test("Broker interrupts should doesn't result in space leaks") { + val br = new Broker[Int] - assert(br.send(123).sync().poll == Some(Return.Unit)) - assert(item.poll == Some(Return(123))) + assert(Offer.select(Offer.const(1), br.recv).poll == Some(Return(1))) + val initial = ObjectSizeCalculator.getObjectSize(br) + + for (_ <- 0 until 1000) { + assert(Offer.select(Offer.const(1), br.recv).poll == Some(Return(1))) + assert(ObjectSizeCalculator.getObjectSize(br) == initial) } } + + test("Broker interrupts should works with orElse") { + val b0, b1 = new Broker[Int] + + val o = b0.recv orElse b1.recv + val f = o.sync() + assert(f.isDefined == false) + + val sendf0 = b0.send(12).sync() + assert(sendf0.isDefined == false) + val sendf1 = b1.send(32).sync() + assert(sendf1.isDefined == true) + assert(f.poll == Some(Return(32))) + + assert(o.sync().poll == Some(Return(12))) + assert(sendf0.poll == Some(Return.Unit)) + } + + test("Broker integrate") { + val br = new Broker[Int] + val offer = Offer.choose(br.recv, Offer.const(999)) + assert(offer.sync().poll == Some(Return(999))) + + val item = br.recv.sync() + assert(item.isDefined == false) + + assert(br.send(123).sync().poll == Some(Return.Unit)) + assert(item.poll == Some(Return(123))) + } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/OfferTest.scala b/util-core/src/test/scala/com/twitter/concurrent/OfferTest.scala index ee5472c59e..1ad41d6755 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/OfferTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/OfferTest.scala @@ -1,366 +1,326 @@ package com.twitter.concurrent -import scala.util.Random - -import org.junit.runner.RunWith +import com.twitter.conversions.DurationOps._ +import com.twitter.util.Await +import com.twitter.util.Future +import com.twitter.util.MockTimer +import com.twitter.util.Promise +import com.twitter.util.Return +import com.twitter.util.Time import org.mockito.Mockito._ -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar - -import com.twitter.conversions.time._ -import com.twitter.util.{Await, Future, MockTimer, Promise, Return, Time} +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.mockito.MockitoSugar +import scala.collection.compat.immutable.LazyList +import scala.util.Random -class SimpleOffer[T](var futures: Stream[Future[Tx[T]]]) extends Offer[T] { - def this(fut: Future[Tx[T]]) = this(Stream.continually(fut)) - def this(tx: Tx[T]) = this(Stream.continually(Future.value(tx))) +class SimpleOffer[T](var futures: LazyList[Future[Tx[T]]]) extends Offer[T] { + def this(fut: Future[Tx[T]]) = this(LazyList.continually(fut)) + def this(tx: Tx[T]) = this(LazyList.continually(Future.value(tx))) - def prepare() = { - val next #:: rest = futures - futures = rest + def prepare(): Future[Tx[T]] = { + val next = futures.head + futures = futures.tail next } } -@RunWith(classOf[JUnitRunner]) -class OfferTest extends WordSpec with MockitoSugar { - import Tx.{Commit, Abort} +class OfferSpecHelper { + val pendingTxs = 0 until 3 map { _ => new Promise[Tx[Int]] } + val offers = pendingTxs map { tx => spy(new SimpleOffer(tx)) } + val offer = Offer.choose(offers: _*) +} - "Offer.map" should { - "apply f in after Tx.ack()" in { - // mockito can't spy on anonymous classes. - val tx = mock[Tx[Int]] - val result = Future.value(Commit(123)) - when(tx.ack()).thenReturn(result) +class TxReadyHelper extends OfferSpecHelper with MockitoSugar { + val tx1 = mock[Tx[Int]] + pendingTxs(1).setValue(tx1) +} - assert(tx.ack() == (result)) - val offer = spy(new SimpleOffer(tx)) +class AllTxsReadyHelper extends OfferSpecHelper with MockitoSugar { + val txs = for (p <- pendingTxs) yield { + val tx = mock[Tx[Int]] + p.setValue(tx) + tx + } +} - val mapped = offer map { i => - (i - 100).toString - } +class OfferOrElseHelper { + val txp = new Promise[Tx[Int]] + val e0 = spy(new SimpleOffer(txp)) +} - val f = mapped.prepare() flatMap { tx => - tx.ack() - } - f match { - case Future(Return(Commit("23"))) => assert(true) - } - } - } +class SyncIntegrationHelper { + val tx2 = new Promise[Tx[Int]] + val e0 = spy( + new SimpleOffer( + Future.value(Tx.aborted: Tx[Int]) #:: (tx2: Future[Tx[Int]]) #:: LazyList.empty + ) + ) + val offer = e0 orElse Offer.const(123) +} - "Offer.choose" should { - class OfferSpecHelper { - val pendingTxs = 0 until 3 map { _ => - new Promise[Tx[Int]] - } - val offers = pendingTxs map { tx => - spy(new SimpleOffer(tx)) - } - val offer = Offer.choose(offers: _*) - } +class OfferTest extends AnyFunSuite with MockitoSugar { - "when a tx is already ready" should { - class TxReadyHelper extends OfferSpecHelper { - val tx1 = mock[Tx[Int]] - pendingTxs(1).setValue(tx1) - } + def await[A](f: Future[A]): A = Await.result(f, 2.seconds) - "prepare() prepares all" in { - val h = new TxReadyHelper - import h._ - - offers foreach { of => - verify(of, never()).prepare() - } - assert(offer.prepare().isDefined == true) - offers foreach { of => - verify(of).prepare() - } - } + import Tx.Commit + import Tx.Abort - "select it" in { - val h = new TxReadyHelper - import h._ + test("Offer.map should apply f in after Tx.ack()") { + // mockito can't spy on anonymous classes. + val tx = mock[Tx[Int]] + val result = Future.value(Commit(123)) + when(tx.ack()).thenReturn(result) - offer.prepare() match { - case Future(Return(tx)) => assert(tx eq tx1) - } - verify(tx1, never()).ack() - verify(tx1, never()).nack() - } + assert(tx.ack() == (result)) + val offer = spy(new SimpleOffer(tx)) - "nack losers" in { - val h = new TxReadyHelper - import h._ - - offer.prepare() - val tx = mock[Tx[Int]] - for (i <- Seq(0, 2)) { - pendingTxs(i).setValue(tx) - } - verify(tx, times(2)).nack() - verify(tx, never()).ack() - } - } + val mapped = offer map { i => (i - 100).toString } - "when a tx is ready after prepare()" should { - "select it" in { - val h = new OfferSpecHelper - import h._ - - val tx = offer.prepare() - assert(tx.isDefined == false) - val tx0 = mock[Tx[Int]] - pendingTxs(0).setValue(tx0) - tx match { - case Future(Return(tx)) => assert(tx eq tx0) - } - } + val f = mapped.prepare() flatMap { tx => tx.ack() } - "nack losers" in { - val h = new OfferSpecHelper - import h._ - - offer.prepare() - pendingTxs(0).setValue(mock[Tx[Int]]) - for (i <- Seq(1, 2)) { - val tx = mock[Tx[Int]] - pendingTxs(i).setValue(tx) - verify(tx).nack() - verify(tx, never()).ack() - } - } - } + assert(await(f) == Commit("23")) + } - "when all txs are ready" should { - class AllTxsReadyHelper extends OfferSpecHelper { - val txs = for (p <- pendingTxs) yield { - val tx = mock[Tx[Int]] - p.setValue(tx) - tx - } - } + test("Offer.choose when a tx is already ready: prepare() prepares all") { + val h = new TxReadyHelper + import h._ - "shuffle winner" in Time.withTimeAt(Time.epoch) { tc => - val h = new AllTxsReadyHelper - import h._ + offers foreach { of => verify(of, never()).prepare() } + assert(offer.prepare().isDefined == true) + offers foreach { of => verify(of).prepare() } + } - val shuffledOffer = Offer.choose(Some(new Random(Time.now.inNanoseconds)), offers) - val histo = new Array[Int](3) - for (_ <- 0 until 1000) { - for (tx <- shuffledOffer.prepare()) - histo(txs.indexOf(tx)) += 1 - } + test("Offer.choose when a tx is already ready should select it") { + val h = new TxReadyHelper + import h._ - assert(histo(0) == 311) - assert(histo(1) == 346) - assert(histo(2) == 343) - } + assert(await(offer.prepare()) == tx1) + verify(tx1, never()).ack() + verify(tx1, never()).nack() + } - "nack losers" in { - val h = new AllTxsReadyHelper - import h._ + test("Offer.choose when a tx is already ready should nack losers") { + val h = new TxReadyHelper + import h.offer + import h.pendingTxs - offer.prepare() match { - case Future(Return(tx)) => - for (loser <- txs if loser ne tx) - verify(loser).nack() - assert(true) - } - } + offer.prepare() + val tx = mock[Tx[Int]] + for (i <- Seq(0, 2)) { + pendingTxs(i).setValue(tx) } + verify(tx, times(2)).nack() + verify(tx, never()).ack() + } - "work with 0 offers" in { - val h = new OfferSpecHelper + test("Offer.choose when a tx is ready after prepare() should select it") { + val h = new OfferSpecHelper + import h._ - val of = Offer.choose() - assert(of.sync().poll == None) + val tx = offer.prepare() + assert(tx.isDefined == false) + val tx0 = mock[Tx[Int]] + pendingTxs(0).setValue(tx0) + + assert(await(tx) == tx0) + } + + test("Offer.choose when a tx is ready after prepare() should nack losers") { + val h = new OfferSpecHelper + import h._ + + offer.prepare() + pendingTxs(0).setValue(mock[Tx[Int]]) + for (i <- Seq(1, 2)) { + val tx = mock[Tx[Int]] + pendingTxs(i).setValue(tx) + verify(tx).nack() + verify(tx, never()).ack() } } - "Offer.sync" should { - "when Tx.prepare is immediately available" should { - "when it commits" in { - val txp = new Promise[Tx[Int]] - val offer = spy(new SimpleOffer(txp)) - val tx = mock[Tx[Int]] - val result = Future.value(Commit(123)) - when(tx.ack()).thenReturn(result) - - assert(tx.ack() == (result)) - txp.setValue(tx) - offer.sync() match { - case Future(Return(123)) => assert(true) - } - verify(tx, times(2)).ack() - verify(tx, never()).nack() - } + test("Offer.choose when all txs are ready should shuffle winner") { + Time.withTimeAt(Time.epoch) { tc => + val h = new AllTxsReadyHelper + import h._ - "retry when it aborts" in { - val txps = new Promise[Tx[Int]] #:: new Promise[Tx[Int]] #:: Stream.empty - val offer = spy(new SimpleOffer(txps)) - val badTx = mock[Tx[Int]] - val result = Future.value(Abort) - when(badTx.ack()).thenReturn(result) + val shuffledOffer = Offer.choose(Some(new Random(Time.now.inNanoseconds)), offers) + val histo = new Array[Int](3) + for (_ <- 0 until 1000) { + for (tx <- shuffledOffer.prepare()) + histo(txs.indexOf(tx)) += 1 + } - assert(badTx.ack() == (result)) - txps(0).setValue(badTx) + assert(histo(0) == 311) + assert(histo(1) == 346) + assert(histo(2) == 343) + } + } - val syncd = offer.sync() + test("Offer.choose when all txs are ready should nack losers") { + val h = new AllTxsReadyHelper + import h._ - assert(syncd.poll == None) - verify(badTx, times(2)).ack() - verify(offer, times(2)).prepare() + val tx = await(offer.prepare()) + for (loser <- txs if loser ne tx) verify(loser).nack() + } - val okTx = mock[Tx[Int]] - val okResult = Future.value(Commit(333)) - when(okTx.ack()).thenReturn(okResult) + test("Offer.choose should work with 0 offers") { + val h = new OfferSpecHelper - assert(okTx.ack() == (okResult)) - txps(1).setValue(okTx) + val of = Offer.choose() + assert(of.sync().poll == None) + } - assert(syncd.poll == Some(Return(333))) + test("Offer.sync when Tx.prepare is immediately available: when it commits") { + val txp = new Promise[Tx[Int]] + val offer = spy(new SimpleOffer(txp)) + val tx = mock[Tx[Int]] + val result = Future.value(Commit(123)) + when(tx.ack()).thenReturn(result) - verify(okTx, times(2)).ack() - verify(offer, times(2)).prepare() - } - } + assert(tx.ack() == (result)) + txp.setValue(tx) - "when Tx.prepare is delayed" should { - "when it commits" in { - val tx = mock[Tx[Int]] - val result = Future.value(Commit(123)) - when(tx.ack()).thenReturn(result) - assert(tx.ack() == (result)) - val offer = spy(new SimpleOffer(tx)) - - offer.sync() match { - case Future(Return(123)) => assert(true) - } - verify(tx, times(2)).ack() - verify(tx, never()).nack() - verify(offer).prepare() - } - } + assert(await(offer.sync()) == 123) + verify(tx, times(2)).ack() + verify(tx, never()).nack() } - "Offer.const" should { - "always provide the same result" in { - val offer = Offer.const(123) + test("Offer.sync when Tx.prepare is immediately available: retry when it aborts") { + val txps = new Promise[Tx[Int]] #:: new Promise[Tx[Int]] #:: LazyList.empty + val offer = spy(new SimpleOffer(txps)) + val badTx = mock[Tx[Int]] + val result = Future.value(Abort) + when(badTx.ack()).thenReturn(result) - assert(offer.sync().poll == Some(Return(123))) - assert(offer.sync().poll == Some(Return(123))) - } + assert(badTx.ack() == (result)) + txps(0).setValue(badTx) - "evaluate argument for each prepare()" in { - var i = 0 - val offer = Offer.const { i = i + 1; i } - assert(offer.sync().poll == Some(Return(1))) - assert(offer.sync().poll == Some(Return(2))) - } - } + val syncd = offer.sync() - "Offer.orElse" should { - "with const orElse" should { - class OfferOrElseHelper { - val txp = new Promise[Tx[Int]] - val e0 = spy(new SimpleOffer(txp)) - } + assert(syncd.poll == None) + verify(badTx, times(2)).ack() + verify(offer, times(2)).prepare() - "prepare orElse event when prepare isn't immediately available" in { - val h = new OfferOrElseHelper - import h._ - - val e1 = Offer.const(123) - val offer = e0 orElse e1 - assert(offer.sync().poll == Some(Return(123))) - verify(e0).prepare() - val tx = mock[Tx[Int]] - txp.setValue(tx) - verify(tx).nack() - verify(tx, never()).ack() - } + val okTx = mock[Tx[Int]] + val okResult = Future.value(Commit(333)) + when(okTx.ack()).thenReturn(okResult) - "not prepare orElse event when the result is immediately available" in { - val h = new OfferOrElseHelper - import h._ - - val e1 = spy(new SimpleOffer(Stream.empty)) - val offer = e0 orElse e1 - val tx = mock[Tx[Int]] - val result = Future.value(Commit(321)) - when(tx.ack()).thenReturn(result) - - assert(tx.ack() == (result)) - txp.setValue(tx) - assert(offer.sync().poll == Some(Return(321))) - verify(e0).prepare() - verify(tx, times(2)).ack() - verify(tx, never()).nack() - verify(e1, never).prepare() - } - } + assert(okTx.ack() == (okResult)) + txps(1).setValue(okTx) - "sync integration: when first transaction aborts" should { - class SyncIntegrationHelper { - val tx2 = new Promise[Tx[Int]] - val e0 = spy( - new SimpleOffer( - Future.value(Tx.aborted: Tx[Int]) #:: (tx2: Future[Tx[Int]]) #:: Stream.empty - ) - ) - val offer = e0 orElse Offer.const(123) - } + assert(syncd.poll == Some(Return(333))) + + verify(okTx, times(2)).ack() + verify(offer, times(2)).prepare() + } - "select first again if tx2 is ready" in { - val h = new SyncIntegrationHelper - import h._ + test("Offer.sync when Tx.prepare is delayed: when it commits") { + val tx = mock[Tx[Int]] + val result = Future.value(Commit(123)) + when(tx.ack()).thenReturn(result) + assert(tx.ack() == (result)) + val offer = spy(new SimpleOffer(tx)) + + assert(await(offer.sync()) == 123) + verify(tx, times(2)).ack() + verify(tx, never()).nack() + verify(offer).prepare() + } - val tx = mock[Tx[Int]] - val result = Future.value(Commit(321)) - when(tx.ack()).thenReturn(result) + test("Offer.const should always provide the same result") { + val offer = Offer.const(123) - assert(tx.ack() == (result)) - tx2.setValue(tx) + assert(offer.sync().poll == Some(Return(123))) + assert(offer.sync().poll == Some(Return(123))) + } - assert(offer.sync().poll == Some(Return(321))) - verify(e0, times(2)).prepare() - verify(tx, times(2)).ack() - verify(tx, never()).nack() - } + test("Offer.const should evaluate argument for each prepare()") { + var i = 0 + val offer = Offer.const { i = i + 1; i } + assert(offer.sync().poll == Some(Return(1))) + assert(offer.sync().poll == Some(Return(2))) + } + + test( + "Offer.orElse with const orElse should prepare orElse event when prepare isn't immediately available") { + val h = new OfferOrElseHelper + import h._ + + val e1 = Offer.const(123) + val offer = e0 orElse e1 + assert(offer.sync().poll == Some(Return(123))) + verify(e0).prepare() + val tx = mock[Tx[Int]] + txp.setValue(tx) + verify(tx).nack() + verify(tx, never()).ack() + } - "select alternative if not ready the second time" in { - val h = new SyncIntegrationHelper - import h._ + test( + "Offer.orElse with const orElse should not prepare orElse event when the result is immediately available") { + val h = new OfferOrElseHelper + import h._ + + val e1 = spy(new SimpleOffer(LazyList.empty)) + val offer = e0 orElse e1 + val tx = mock[Tx[Int]] + val result = Future.value(Commit(321)) + when(tx.ack()).thenReturn(result) + + assert(tx.ack() == (result)) + txp.setValue(tx) + assert(offer.sync().poll == Some(Return(321))) + verify(e0).prepare() + verify(tx, times(2)).ack() + verify(tx, never()).nack() + verify(e1, never).prepare() + } + test( + "Offer.orElse sync integration: when first transaction aborts should select first again if tx2 is ready") { + val h = new SyncIntegrationHelper + import h._ + + val tx = mock[Tx[Int]] + val result = Future.value(Commit(321)) + when(tx.ack()).thenReturn(result) + + assert(tx.ack() == (result)) + tx2.setValue(tx) + + assert(offer.sync().poll == Some(Return(321))) + verify(e0, times(2)).prepare() + verify(tx, times(2)).ack() + verify(tx, never()).nack() + } - assert(offer.sync().poll == Some(Return(123))) - verify(e0, times(2)).prepare() + test( + "Offer.orElse sync integration: when first transaction aborts should select alternative if not ready the second time") { + val h = new SyncIntegrationHelper + import h._ - val tx = mock[Tx[Int]] - tx2.setValue(tx) - verify(tx).nack() - } - } + assert(offer.sync().poll == Some(Return(123))) + verify(e0, times(2)).prepare() + + val tx = mock[Tx[Int]] + tx2.setValue(tx) + verify(tx).nack() } - "Offer.foreach" should { - "synchronize on offers forever" in { - val b = new Broker[Int] - var count = 0 - b.recv foreach { _ => - count += 1 - } - assert(count == 0) - assert(b.send(1).sync().isDefined == true) - assert(count == 1) - assert(b.send(1).sync().isDefined == true) - assert(count == 2) - } + test("Offer.foreach should synchronize on offers forever") { + val b = new Broker[Int] + var count = 0 + b.recv foreach { _ => count += 1 } + assert(count == 0) + assert(b.send(1).sync().isDefined == true) + assert(count == 1) + assert(b.send(1).sync().isDefined == true) + assert(count == 2) } - "Offer.timeout" should { - "be available after timeout (prepare)" in Time.withTimeAt(Time.epoch) { tc => + test("Offer.timeout should be available after timeout (prepare)") { + Time.withTimeAt(Time.epoch) { tc => implicit val timer = new MockTimer val e = Offer.timeout(10.seconds) assert(e.prepare().isDefined == false) @@ -371,15 +331,13 @@ class OfferTest extends WordSpec with MockitoSugar { timer.tick() assert(e.prepare().isDefined == true) } + } - "cancel timer tasks when losing" in Time.withTimeAt(Time.epoch) { tc => + test("Offer.timeout should cancel timer tasks when losing") { + Time.withTimeAt(Time.epoch) { tc => implicit val timer = new MockTimer - val e10 = Offer.timeout(10.seconds) map { _ => - 10 - } - val e5 = Offer.timeout(5.seconds) map { _ => - 5 - } + val e10 = Offer.timeout(10.seconds) map { _ => 10 } + val e5 = Offer.timeout(5.seconds) map { _ => 5 } val item = Offer.select(e5, e10) assert(item.poll == None) @@ -395,47 +353,43 @@ class OfferTest extends WordSpec with MockitoSugar { } } - "Offer.prioritize" should { - "prioritize high priority offer" in { - val offers = 0 until 10 map { i => - val tx = mock[Tx[Int]] - val result = Future.value(Commit(i)) - when(tx.ack()).thenReturn(result) - new SimpleOffer(Future.value(tx)) - } - val chosenOffer = Offer.prioritize(offers: _*) - val of = chosenOffer.sync() - assert(of.isDefined == true) - assert(Await.result(of) == 0) + test("Offer.prioritize should prioritize high priority offer") { + val offers = 0 until 10 map { i => + val tx = mock[Tx[Int]] + val result = Future.value(Commit(i)) + when(tx.ack()).thenReturn(result) + new SimpleOffer(Future.value(tx)) } + val chosenOffer = Offer.prioritize(offers: _*) + val of = chosenOffer.sync() + assert(of.isDefined == true) + assert(Await.result(of) == 0) } - "Integration" should { - "select across multiple brokers" in { - val b0 = new Broker[Int] - val b1 = new Broker[String] - - val o = Offer.choose( - b0.send(123) const { "put!" }, - b1.recv - ) - - val f = o.sync() - assert(f.isDefined == false) - assert(b1.send("hey").sync().isDefined == true) - assert(f.isDefined == true) - assert(Await.result(f) == "hey") - - val gf = b0.recv.sync() - assert(gf.isDefined == false) - val of = o.sync() - assert(of.isDefined == true) - assert(Await.result(of) == "put!") - assert(gf.isDefined == true) - assert(Await.result(gf) == 123) - - // syncing again fails. - assert(o.sync().isDefined == false) - } + test("Integration should select across multiple brokers") { + val b0 = new Broker[Int] + val b1 = new Broker[String] + + val o = Offer.choose( + b0.send(123) const { "put!" }, + b1.recv + ) + + val f = o.sync() + assert(f.isDefined == false) + assert(b1.send("hey").sync().isDefined == true) + assert(f.isDefined == true) + assert(Await.result(f) == "hey") + + val gf = b0.recv.sync() + assert(gf.isDefined == false) + val of = o.sync() + assert(of.isDefined == true) + assert(Await.result(of) == "put!") + assert(gf.isDefined == true) + assert(Await.result(gf) == 123) + + // syncing again fails. + assert(o.sync().isDefined == false) } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/OnceTest.scala b/util-core/src/test/scala/com/twitter/concurrent/OnceTest.scala index 0371ade6e7..afca463483 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/OnceTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/OnceTest.scala @@ -1,13 +1,10 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Promise, Await} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class OnceTest extends FunSuite { +class OnceTest extends AnyFunSuite { test("Once.apply should only be applied once") { var x = 0 val once = Once { x += 1 } diff --git a/util-core/src/test/scala/com/twitter/concurrent/PeriodTest.scala b/util-core/src/test/scala/com/twitter/concurrent/PeriodTest.scala index 8161b9b749..2e0822040e 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/PeriodTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/PeriodTest.scala @@ -1,12 +1,9 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class PeriodTest extends FunSuite { +class PeriodTest extends AnyFunSuite { test("Period#numPeriods should behave reasonably") { val period = new Period(1.second) val dur = 1.second diff --git a/util-core/src/test/scala/com/twitter/concurrent/SchedulerTest.scala b/util-core/src/test/scala/com/twitter/concurrent/SchedulerTest.scala index 778d5cae30..523c83b92a 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/SchedulerTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/SchedulerTest.scala @@ -1,22 +1,21 @@ package com.twitter.concurrent -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util._ import java.util.concurrent.{Executors, TimeUnit} import java.util.logging.{Handler, Level, LogRecord} -import org.junit.runner.RunWith -import org.scalatest.{FunSuite, Matchers} import org.scalatest.concurrent.{Eventually, IntegrationPatience} -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers -abstract class LocalSchedulerTest(lifo: Boolean) extends FunSuite with Matchers { +abstract class LocalSchedulerTest(lifo: Boolean) extends AnyFunSuite with Matchers { private val scheduler = new LocalScheduler(lifo) def submit(f: => Unit): Unit = scheduler.submit(new Runnable { def run(): Unit = f }) - val N = 100 + val N: Int = 100 test("run the first submitter immediately") { var ok = false @@ -151,17 +150,13 @@ abstract class LocalSchedulerTest(lifo: Boolean) extends FunSuite with Matchers val record = sampleBlockingFraction(0.0) assert(record == null) } - } -@RunWith(classOf[JUnitRunner]) class LocalSchedulerFifoTest extends LocalSchedulerTest(false) -@RunWith(classOf[JUnitRunner]) class LocalSchedulerLifoTest extends LocalSchedulerTest(true) -@RunWith(classOf[JUnitRunner]) -class ThreadPoolSchedulerTest extends FunSuite with Eventually with IntegrationPatience { +class ThreadPoolSchedulerTest extends AnyFunSuite with Eventually with IntegrationPatience { test("works") { val p = new Promise[Unit] val scheduler = new ThreadPoolScheduler("test") @@ -169,8 +164,7 @@ class ThreadPoolSchedulerTest extends FunSuite with Eventually with IntegrationP def run(): Unit = p.setDone() }) - eventually { p.isDone } - + Await.result(p, 5.seconds) scheduler.shutdown() } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/SerializedTest.scala b/util-core/src/test/scala/com/twitter/concurrent/SerializedTest.scala index 7279bdffcc..2fb27928c8 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/SerializedTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/SerializedTest.scala @@ -1,46 +1,41 @@ package com.twitter.concurrent import java.util.concurrent.CountDownLatch +import org.scalatest.funsuite.AnyFunSuite -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) -class SerializedTest extends WordSpec with Serialized { - "Serialized" should { - "runs blocks, one at a time, in the order received" in { - val t1CallsSerializedFirst = new CountDownLatch(1) - val t1FinishesWork = new CountDownLatch(1) - val orderOfExecution = new collection.mutable.ListBuffer[Thread] +class SerializedTest extends AnyFunSuite with Serialized { + test("Serialized runs blocks, one at a time, in the order received") { + val t1CallsSerializedFirst = new CountDownLatch(1) + val t1FinishesWork = new CountDownLatch(1) + val orderOfExecution = new collection.mutable.ListBuffer[Thread] - val t1 = new Thread { - override def run: Unit = { - serialized { - t1CallsSerializedFirst.countDown() - t1FinishesWork.await() - orderOfExecution += this - () - } + val t1 = new Thread { + override def run: Unit = { + serialized { + t1CallsSerializedFirst.countDown() + t1FinishesWork.await() + orderOfExecution += this + () } } + } - val t2 = new Thread { - override def run: Unit = { - t1CallsSerializedFirst.await() - serialized { - orderOfExecution += this - () - } - t1FinishesWork.countDown() + val t2 = new Thread { + override def run: Unit = { + t1CallsSerializedFirst.await() + serialized { + orderOfExecution += this + () } + t1FinishesWork.countDown() } + } - t1.start() - t2.start() - t1.join() - t2.join() + t1.start() + t2.start() + t1.join() + t2.join() - assert(orderOfExecution.toList == List(t1, t2)) - } + assert(orderOfExecution.toList == List(t1, t2)) } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/SpoolSourceTest.scala b/util-core/src/test/scala/com/twitter/concurrent/SpoolSourceTest.scala index 942b674255..70b29f584a 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/SpoolSourceTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/SpoolSourceTest.scala @@ -1,113 +1,109 @@ package com.twitter.concurrent -import com.twitter.util.{Return, Await} +import com.twitter.util.Await +import com.twitter.util.Return import java.util.concurrent.atomic.AtomicInteger -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class SpoolSourceTest extends WordSpec { - "SpoolSource" should { - class SpoolSourceHelper { - val source = new SpoolSource[Int] - } +import org.scalatest.funsuite.AnyFunSuite - "add values to the spool, ignoring values after close" in { - val h = new SpoolSourceHelper - import h._ - - val futureSpool = source() - source.offer(1) - source.offer(2) - source.offer(3) - source.close() - source.offer(4) - source.offer(5) - assert(Await.result(futureSpool flatMap (_.toSeq)) == Seq(1, 2, 3)) - } +class SpoolSourceTest extends AnyFunSuite { + class SpoolSourceHelper { + val source = new SpoolSource[Int] + } - "add values to the spool, ignoring values after offerAndClose" in { - val h = new SpoolSourceHelper - import h._ - - val futureSpool = source() - source.offer(1) - source.offer(2) - source.offerAndClose(3) - source.offer(4) - source.offer(5) - assert(Await.result(futureSpool flatMap (_.toSeq)) == Seq(1, 2, 3)) - assert(Await.result(source.closed.liftToTry) == Return(())) - } + test("SpoolSource should add values to the spool, ignoring values after close") { + val h = new SpoolSourceHelper + import h._ + + val futureSpool = source() + source.offer(1) + source.offer(2) + source.offer(3) + source.close() + source.offer(4) + source.offer(5) + assert(Await.result(futureSpool flatMap (_.toSeq)) == Seq(1, 2, 3)) + } - "return multiple Future Spools that only see values added later" in { - val h = new SpoolSourceHelper - import h._ - - val futureSpool1 = source() - source.offer(1) - val futureSpool2 = source() - source.offer(2) - val futureSpool3 = source() - source.offer(3) - val futureSpool4 = source() - source.close() - assert(Await.result(futureSpool1 flatMap (_.toSeq)) == Seq(1, 2, 3)) - assert(Await.result(futureSpool2 flatMap (_.toSeq)) == Seq(2, 3)) - assert(Await.result(futureSpool3 flatMap (_.toSeq)) == Seq(3)) - assert(Await.result(futureSpool4).isEmpty == true) - } + test("SpoolSource should add values to the spool, ignoring values after offerAndClose") { + val h = new SpoolSourceHelper + import h._ + + val futureSpool = source() + source.offer(1) + source.offer(2) + source.offerAndClose(3) + source.offer(4) + source.offer(5) + assert(Await.result(futureSpool flatMap (_.toSeq)) == Seq(1, 2, 3)) + assert(Await.result(source.closed.liftToTry) == Return(())) + } + + test("SpoolSource should return multiple Future Spools that only see values added later") { + val h = new SpoolSourceHelper + import h._ + + val futureSpool1 = source() + source.offer(1) + val futureSpool2 = source() + source.offer(2) + val futureSpool3 = source() + source.offer(3) + val futureSpool4 = source() + source.close() + assert(Await.result(futureSpool1 flatMap (_.toSeq)) == Seq(1, 2, 3)) + assert(Await.result(futureSpool2 flatMap (_.toSeq)) == Seq(2, 3)) + assert(Await.result(futureSpool3 flatMap (_.toSeq)) == Seq(3)) + assert(Await.result(futureSpool4).isEmpty == true) + } - "throw exception and close spool when exception is raised" in { - val h = new SpoolSourceHelper - import h._ - - val futureSpool1 = source() - source.offer(1) - source.raise(new Exception("sad panda")) - val futureSpool2 = source() - source.offer(1) - intercept[Exception] { - Await.result(futureSpool1 flatMap (_.toSeq)) - } - assert(Await.result(futureSpool2).isEmpty == true) + test("SpoolSource should throw exception and close spool when exception is raised") { + val h = new SpoolSourceHelper + import h._ + + val futureSpool1 = source() + source.offer(1) + source.raise(new Exception("sad panda")) + val futureSpool2 = source() + source.offer(1) + intercept[Exception] { + Await.result(futureSpool1 flatMap (_.toSeq)) } + assert(Await.result(futureSpool2).isEmpty == true) + } - "invoke registered interrupt handlers" in { - val h = new SpoolSourceHelper + test("SpoolSource should invoke registered interrupt handlers") { + val h = new SpoolSourceHelper - val fired = new AtomicInteger(0) + val fired = new AtomicInteger(0) - val isource = new SpoolSource[Int]({ - case exc => - fired.addAndGet(1) - }) + val isource = new SpoolSource[Int]({ + case exc => + fired.addAndGet(1) + }) - val a = isource() - val b = a.flatMap(_.tail) - val c = b.flatMap(_.tail) - val d = c.flatMap(_.tail) + val a = isource() + val b = a.flatMap(_.tail) + val c = b.flatMap(_.tail) + val d = c.flatMap(_.tail) - assert(fired.get() == 0) + assert(fired.get() == 0) - a.raise(new Exception("sad panda 0")) - assert(fired.get() == 1) + a.raise(new Exception("sad panda 0")) + assert(fired.get() == 1) - isource.offer(1) - b.raise(new Exception("sad panda 1")) - assert(fired.get() == 2) + isource.offer(1) + b.raise(new Exception("sad panda 1")) + assert(fired.get() == 2) - isource.offer(2) - c.raise(new Exception("sad panda 2")) - assert(fired.get() == 3) + isource.offer(2) + c.raise(new Exception("sad panda 2")) + assert(fired.get() == 3) - // ignore exceptions once the source is closed - isource.offerAndClose(3) - d.raise(new Exception("sad panda 3")) - assert(fired.get() == 3) + // ignore exceptions once the source is closed + isource.offerAndClose(3) + d.raise(new Exception("sad panda 3")) + assert(fired.get() == 3) - assert(Await.result(a flatMap (_.toSeq)) == Seq(1, 2, 3)) - } + assert(Await.result(a flatMap (_.toSeq)) == Seq(1, 2, 3)) } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/SpoolTest.scala b/util-core/src/test/scala/com/twitter/concurrent/SpoolTest.scala index f51abd6697..dfa743153a 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/SpoolTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/SpoolTest.scala @@ -1,495 +1,464 @@ package com.twitter.concurrent -import com.twitter.concurrent.Spool.{*::, seqToSpool} -import com.twitter.conversions.time.intToTimeableNumber -import com.twitter.util.{Await, Future, Promise, Return, Throw} +import com.twitter.concurrent.Spool.*:: +import com.twitter.concurrent.Spool.seqToSpool +import com.twitter.conversions.DurationOps._ +import com.twitter.util.Await +import com.twitter.util.Future +import com.twitter.util.Promise +import com.twitter.util.Return +import com.twitter.util.Throw import java.io.EOFException -import org.junit.runner.RunWith import org.scalacheck.Arbitrary -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks import scala.collection.mutable.ArrayBuffer -@RunWith(classOf[JUnitRunner]) -class SpoolTest extends WordSpec with GeneratorDrivenPropertyChecks { - "Empty Spool" should { - val s = Spool.empty[Int] +class SpoolTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { - "iterate over all elements" in { - val xs = new ArrayBuffer[Int] - s foreach { xs += _ } - assert(xs.size == 0) - } + def await[A](f: Future[A]): A = Await.result(f, 2.seconds) + val resolvedSpoolFuture = 1 *:: Future.value(2 *:: Future.value(Spool.empty[Int])) + val emptySpool = Spool.empty[Int] - "map" in { - assert((s map { _ * 2 }) == Spool.empty[Int]) - } + class SimpleDelayedSpoolHelper { + val p = new Promise[Spool[Int]] + val p1 = new Promise[Spool[Int]] + val p2 = new Promise[Spool[Int]] + val s = 1 *:: p + } - "mapFuture" in { - val mapFuture = s mapFuture { Future.value(_) } - assert(mapFuture.poll == Some(Return(s))) - } + /** + * Confirms that the given operation does not consume an entire Spool, and then + * returns the resulting Spool and tail check for further validation. + */ + def applyLazily( + f: Spool[Int] => Future[Spool[Int]] + ): (Future[Spool[Int]], Future[Spool[Int]]) = { + val tail = new Promise[Spool[Int]] + + // A spool where only the head is valid. + def spool: Spool[Int] = 0 *:: { tail.setException(new Exception); tail } + + // create, apply, poll + val s = f(spool) + assert(!tail.isDefined) + (s, tail) + } - "deconstruct" in { - assert(s match { - case x *:: Future(rest) => false - case _ => true - }) - } + test("Empty Spool should iterate over all elements") { + val xs = new ArrayBuffer[Int] + emptySpool foreach { xs += _ } + assert(xs.size == 0) + } - "append via ++" in { - assert((s ++ Spool.empty[Int]) == Spool.empty[Int]) - assert((Spool.empty[Int] ++ s) == Spool.empty[Int]) + test("Empty Spool should map") { + assert((emptySpool map { _ * 2 }) == Spool.empty[Int]) + } - val s2 = s ++ (3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) - assert(Await.result(s2.toSeq) == Seq(3, 4)) - } + test("Empty Spool should mapFuture") { + val mapFuture = emptySpool mapFuture { Future.value(_) } + assert(mapFuture.poll == Some(Return(emptySpool))) + } - "append via ++ with Future rhs" in { - assert(Await.result(s ++ Future(Spool.empty[Int])) == Spool.empty[Int]) - assert(Await.result(Spool.empty[Int] ++ Future(s)) == Spool.empty[Int]) + test("Empty Spool should deconstruct") { + assert(emptySpool match { + case x *:: _ => false + case _ => true + }) + } - val s2 = s ++ Future(3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) - assert(Await.result(s2 flatMap (_.toSeq)) == Seq(3, 4)) - } + test("Empty Spool should append via ++") { + assert((emptySpool ++ Spool.empty[Int]) == Spool.empty[Int]) + assert((Spool.empty[Int] ++ emptySpool) == Spool.empty[Int]) - "flatMap" in { - val f = (x: Int) => - Future(x.toString *:: Future.value((x * 2).toString *:: Future.value(Spool.empty[String]))) - assert(Await.result(s flatMap f) == Spool.empty[Int]) - } + val s2 = emptySpool ++ (3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) + assert(await(s2.toSeq) == Seq(3, 4)) + } - "fold left" in { - val fold = s.foldLeft(0) { (x, y) => - x + y - } - assert(Await.result(fold) == 0) - } + test("Empty Spool should append via ++ with Future rhs") { + assert(await(emptySpool ++ Future(Spool.empty[Int])) == Spool.empty[Int]) + assert(await(Spool.empty[Int] ++ Future(emptySpool)) == Spool.empty[Int]) - "reduce left" in { - val fold = s.reduceLeft { (x, y) => - x + y - } - intercept[UnsupportedOperationException] { - Await.result(fold) - } - } + val s2 = emptySpool ++ Future(3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) + assert(await(s2 flatMap (_.toSeq)) == Seq(3, 4)) + } - "zip with empty" in { - val result = s.zip(Spool.empty[Int]) - assert(Await.result(result.toSeq) == Nil) - } + test("Empty Spool should flatMap") { + val f = (x: Int) => + Future(x.toString *:: Future.value((x * 2).toString *:: Future.value(Spool.empty[String]))) + assert(await(emptySpool flatMap f) == Spool.empty[Int]) + } - "zip with non-empty" in { - val result = s.zip(Seq(1, 2, 3).toSpool) - assert(Await.result(result.toSeq) == Nil) - } + test("Empty Spool should fold left") { + val fold = emptySpool.foldLeft(0) { (x, y) => x + y } + assert(await(fold) == 0) + } - "take" in { - assert(s.take(10) == Spool.empty[Int]) + test("Empty Spool should reduce left") { + val fold = emptySpool.reduceLeft { (x, y) => x + y } + intercept[UnsupportedOperationException] { + await(fold) } } - "Simple resolved Spool" should { - val s = 1 *:: Future.value(2 *:: Future.value(Spool.empty[Int])) + test("Empty Spool should zip with empty") { + val result = emptySpool.zip(Spool.empty[Int]) + assert(await(result.toSeq) == Nil) + } - "iterate over all elements" in { - val xs = new ArrayBuffer[Int] - s foreach { xs += _ } - assert(xs.toSeq == Seq(1, 2)) - } + test("Empty Spool should zip with non-empty") { + val result = emptySpool.zip(Seq(1, 2, 3).toSpool) + assert(await(result.toSeq) == Nil) + } - "buffer to a sequence" in { - assert(Await.result(s.toSeq) == Seq(1, 2)) - } + test("Empty Spool should take") { + assert(emptySpool.take(10) == Spool.empty[Int]) + } - "map" in { - assert(Await.result(s.map { _ * 2 }.toSeq) == Seq(2, 4)) - } + test("Simple resolved spool should iterate over all elements") { + val xs = new ArrayBuffer[Int] + resolvedSpoolFuture foreach { xs += _ } + assert(xs.toSeq == Seq(1, 2)) + } - "mapFuture" in { - val f = s.mapFuture { Future.value(_) }.flatMap { _.toSeq }.poll - assert(f == Some(Return(Seq(1, 2)))) - } + test("Simple resolved spool should buffer to a sequence") { + assert(await(resolvedSpoolFuture.toSeq) == Seq(1, 2)) + } - "deconstruct" in { - assert(s match { - case x *:: Future(Return(rest)) => - assert(x == 1) - rest match { - case y *:: Future(Return(rest)) if y == 2 && rest.isEmpty => true - } - }) - } + test("Simple resolved spool should map") { + assert(await(resolvedSpoolFuture.map { _ * 2 }.toSeq) == Seq(2, 4)) + } - "append via ++" in { - assert(Await.result((s ++ Spool.empty[Int]).toSeq) == Seq(1, 2)) - assert(Await.result((Spool.empty[Int] ++ s).toSeq) == Seq(1, 2)) + test("Simple resolved spool should mapFuture") { + val f = resolvedSpoolFuture.mapFuture { Future.value(_) }.flatMap { _.toSeq }.poll + assert(f == Some(Return(Seq(1, 2)))) + } - val s2 = s ++ (3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) - assert(Await.result(s2.toSeq) == Seq(1, 2, 3, 4)) - } + test("Simple resolved spool should deconstruct") { + assert(resolvedSpoolFuture match { + case x *:: rest => + assert(x == 1) + await(rest) match { + case y *:: rest if y == 2 => await(rest).isEmpty + case _ => fail("expected 3 elements in spool") + } + case _ => fail("expected 3 elements in spool") + }) + } - "append via ++ with Future rhs" in { - assert(Await.result(s ++ Future.value(Spool.empty[Int]) flatMap (_.toSeq)) == Seq(1, 2)) - assert(Await.result(Spool.empty[Int] ++ Future.value(s) flatMap (_.toSeq)) == Seq(1, 2)) + test("Simple resolved spool should append via ++") { + assert(await((resolvedSpoolFuture ++ Spool.empty[Int]).toSeq) == Seq(1, 2)) + assert(await((Spool.empty[Int] ++ resolvedSpoolFuture).toSeq) == Seq(1, 2)) - val s2 = s ++ Future.value(3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) - assert(Await.result(s2 flatMap (_.toSeq)) == Seq(1, 2, 3, 4)) - } + val s2 = resolvedSpoolFuture ++ (3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) + assert(await(s2.toSeq) == Seq(1, 2, 3, 4)) + } - "flatMap" in { - val f = (x: Int) => Future(Seq(x.toString, (x * 2).toString).toSpool) - val s2 = s flatMap f - assert(Await.result(s2 flatMap (_.toSeq)) == Seq("1", "2", "2", "4")) - } + test("Simple resolved spool should append via ++ with Future rhs") { + assert( + await(resolvedSpoolFuture ++ Future.value(Spool.empty[Int]) flatMap (_.toSeq)) == Seq(1, 2)) + assert( + await(Spool.empty[Int] ++ Future.value(resolvedSpoolFuture) flatMap (_.toSeq)) == Seq(1, 2)) - "fold left" in { - val fold = s.foldLeft(0) { (x, y) => - x + y - } - assert(Await.result(fold) == 3) - } + val s2 = + resolvedSpoolFuture ++ Future.value(3 *:: Future.value(4 *:: Future.value(Spool.empty[Int]))) + assert(await(s2 flatMap (_.toSeq)) == Seq(1, 2, 3, 4)) + } - "reduce left" in { - val fold = s.reduceLeft { (x, y) => - x + y - } - assert(Await.result(fold) == 3) - } + test("Simple resolved spool should flatMap") { + val f = (x: Int) => Future(Seq(x.toString, (x * 2).toString).toSpool) + val s2 = resolvedSpoolFuture flatMap f + assert(await(s2 flatMap (_.toSeq)) == Seq("1", "2", "2", "4")) + } - "zip with empty" in { - val zip = s.zip(Spool.empty[Int]) - assert(Await.result(zip.toSeq) == Nil) - } + test("Simple resolved spool should fold left") { + val fold = resolvedSpoolFuture.foldLeft(0) { (x, y) => x + y } + assert(await(fold) == 3) + } - "zip with same size spool" in { - val zip = s.zip(Seq("a", "b").toSpool) - assert(Await.result(zip.toSeq) == Seq((1, "a"), (2, "b"))) - } + test("Simple resolved spool should reduce left") { + val fold = resolvedSpoolFuture.reduceLeft { (x, y) => x + y } + assert(await(fold) == 3) + } - "zip with larger spool" in { - val zip = s.zip(Seq("a", "b", "c", "d").toSpool) - assert(Await.result(zip.toSeq) == Seq((1, "a"), (2, "b"))) - } + test("Simple resolved spool should zip with empty") { + val zip = resolvedSpoolFuture.zip(Spool.empty[Int]) + assert(await(zip.toSeq) == Nil) + } - "be roundtrippable through toSeq/toSpool" in { - val seq = (0 to 10).toSeq - assert(Await.result(seq.toSpool.toSeq) == seq) - } + test("Simple resolved spool should zip with same size spool") { + val zip = resolvedSpoolFuture.zip(Seq("a", "b").toSpool) + assert(await(zip.toSeq) == Seq((1, "a"), (2, "b"))) + } + + test("Simple resolved spool should zip with larger spool") { + val zip = resolvedSpoolFuture.zip(Seq("a", "b", "c", "d").toSpool) + assert(await(zip.toSeq) == Seq((1, "a"), (2, "b"))) + } - "flatten via flatMap of toSpool" in { - val spool = Seq(Seq(1, 2), Seq(3, 4)).toSpool - val seq = Await.result(spool.toSeq) + test("Simple resolved spool should be roundtrippable through toSeq/toSpool") { + val seq = (0 to 10).toSeq + assert(await(seq.toSpool.toSeq) == seq) + } - val flatSpool = - spool.flatMap { inner => - Future.value(inner.toSpool) - } + test("Simple resolved spool should flatten via flatMap of toSpool") { + val spool = Seq(Seq(1, 2), Seq(3, 4)).toSpool + val seq = await(spool.toSeq) - assert(Await.result(flatSpool.flatMap(_.toSeq)) == seq.flatten) - } + val flatSpool = + spool.flatMap { inner => Future.value(inner.toSpool) } - "take" in { - val ls = (1 to 4).toSeq.toSpool - assert(Await.result(ls.take(2).toSeq) == Seq(1, 2)) - assert(Await.result(ls.take(1).toSeq) == Seq(1)) - assert(Await.result(ls.take(0).toSeq) == Seq.empty) - assert(Await.result(ls.take(-2).toSeq) == Seq.empty) - } + assert(await(flatSpool.flatMap(_.toSeq)) == seq.flatten) + } + + test("Simple resolved spool should take") { + val ls = (1 to 4).toSeq.toSpool + assert(await(ls.take(2).toSeq) == Seq(1, 2)) + assert(await(ls.take(1).toSeq) == Seq(1)) + assert(await(ls.take(0).toSeq) == Seq.empty) + assert(await(ls.take(-2).toSeq) == Seq.empty) } - "Simple resolved spool with EOFException" should { + test("Simple resolved spool with EOF iteration on EOFException") { val p = new Promise[Spool[Int]](Throw(new EOFException("sad panda"))) val s = 1 *:: Future.value(2 *:: p) - - "EOF iteration on EOFException" in { - val xs = new ArrayBuffer[Option[Int]] - s foreachElem { xs += _ } - assert(xs.toSeq == Seq(Some(1), Some(2), None)) - } + val xs = new ArrayBuffer[Option[Int]] + s foreachElem { xs += _ } + assert(xs.toSeq == Seq(Some(1), Some(2), None)) } - "Simple resolved spool with error" should { + test("Simple resolved spool with error should return with exception on error") { val p = new Promise[Spool[Int]](Throw(new Exception("sad panda"))) val s = 1 *:: Future.value(2 *:: p) - - "return with exception on error" in { - val xs = new ArrayBuffer[Option[Int]] - s foreachElem { xs += _ } - intercept[Exception] { - Await.result(s.toSeq) - } - } - - "return with exception on error in callback" in { - val xs = new ArrayBuffer[Option[Int]] - val f = s foreach { _ => - throw new Exception("sad panda") - } - intercept[Exception] { - Await.result(f) - } + val xs = new ArrayBuffer[Option[Int]] + s foreachElem { xs += _ } + intercept[Exception] { + await(s.toSeq) } + } - "return with exception on EOFException in callback" in { - val xs = new ArrayBuffer[Option[Int]] - val f = s foreach { _ => - throw new EOFException("sad panda") - } - intercept[EOFException] { - Await.result(f) - } + test("Simple resolved spool with error should return with exception on error in callback") { + val p = new Promise[Spool[Int]](Throw(new Exception("sad panda"))) + val s = 1 *:: Future.value(2 *:: p) + val xs = new ArrayBuffer[Option[Int]] + val f = s foreach { _ => throw new Exception("sad panda") } + intercept[Exception] { + await(f) } } - "Simple delayed Spool" should { - class SimpleDelayedSpoolHelper { - val p = new Promise[Spool[Int]] - val p1 = new Promise[Spool[Int]] - val p2 = new Promise[Spool[Int]] - val s = 1 *:: p + test( + "Simple resolved spool with error should return with exception on EOFException in callback") { + val p = new Promise[Spool[Int]](Throw(new Exception("sad panda"))) + val s = 1 *:: Future.value(2 *:: p) + val xs = new ArrayBuffer[Option[Int]] + val f = s foreach { _ => throw new EOFException("sad panda") } + intercept[EOFException] { + await(f) } + } - "iterate as results become available" in { - val h = new SimpleDelayedSpoolHelper - import h._ - - val xs = new ArrayBuffer[Int] - s foreach { xs += _ } - assert(xs.toSeq == Seq(1)) - p() = Return(2 *:: p1) - assert(xs.toSeq == Seq(1, 2)) - p1() = Return(Spool.empty) - assert(xs.toSeq == Seq(1, 2)) - } + test("Simple delayed Spool should iterate as results become available") { + val h = new SimpleDelayedSpoolHelper + import h._ + + val xs = new ArrayBuffer[Int] + s foreach { xs += _ } + assert(xs.toSeq == Seq(1)) + p() = Return(2 *:: p1) + assert(xs.toSeq == Seq(1, 2)) + p1() = Return(Spool.empty) + assert(xs.toSeq == Seq(1, 2)) + } - "EOF iteration on EOFException" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should EOF iteration on EOFException") { + val h = new SimpleDelayedSpoolHelper + import h._ - val xs = new ArrayBuffer[Option[Int]] - s foreachElem { xs += _ } - assert(xs.toSeq == Seq(Some(1))) - p() = Throw(new EOFException("sad panda")) - assert(xs.toSeq == Seq(Some(1), None)) - } + val xs = new ArrayBuffer[Option[Int]] + s foreachElem { xs += _ } + assert(xs.toSeq == Seq(Some(1))) + p() = Throw(new EOFException("sad panda")) + assert(xs.toSeq == Seq(Some(1), None)) + } - "return with exception on error" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should return with exception on error") { + val h = new SimpleDelayedSpoolHelper + import h._ - val xs = new ArrayBuffer[Option[Int]] - s foreachElem { xs += _ } - assert(xs.toSeq == Seq(Some(1))) - p() = Throw(new Exception("sad panda")) - intercept[Exception] { - Await.result(s.toSeq) - } + val xs = new ArrayBuffer[Option[Int]] + s foreachElem { xs += _ } + assert(xs.toSeq == Seq(Some(1))) + p() = Throw(new Exception("sad panda")) + intercept[Exception] { + await(s.toSeq) } + } - "return with exception on error in callback" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should return with exception on error in callback") { + val h = new SimpleDelayedSpoolHelper + import h._ - val xs = new ArrayBuffer[Option[Int]] - val f = s foreach { _ => - throw new Exception("sad panda") - } - p() = Return(2 *:: p1) - intercept[Exception] { - Await.result(f) - } + val xs = new ArrayBuffer[Option[Int]] + val f = s foreach { _ => throw new Exception("sad panda") } + p() = Return(2 *:: p1) + intercept[Exception] { + await(f) } + } - "return with exception on EOFException in callback" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should return with exception on EOFException in callback") { + val h = new SimpleDelayedSpoolHelper + import h._ - val xs = new ArrayBuffer[Option[Int]] - val f = s foreach { _ => - throw new EOFException("sad panda") - } - p() = Return(2 *:: p1) - intercept[EOFException] { - Await.result(f) - } + val xs = new ArrayBuffer[Option[Int]] + val f = s foreach { _ => throw new EOFException("sad panda") } + p() = Return(2 *:: p1) + intercept[EOFException] { + await(f) } + } - "return a buffered seq when complete" in { - val h = new SimpleDelayedSpoolHelper - import h._ - - val f = s.toSeq - assert(f.isDefined == false) - p() = Return(2 *:: p1) - assert(f.isDefined == false) - p1() = Return(Spool.empty) - assert(f.isDefined == true) - assert(Await.result(f) == Seq(1, 2)) - } - - "deconstruct" in { - val h = new SimpleDelayedSpoolHelper - import h._ - - assert(s match { - case fst *:: rest if fst == 1 && !rest.isDefined => true - case _ => false - }) - } + test("Simple delayed Spool should return a buffered seq when complete") { + val h = new SimpleDelayedSpoolHelper + import h._ + + val f = s.toSeq + assert(f.isDefined == false) + p() = Return(2 *:: p1) + assert(f.isDefined == false) + p1() = Return(Spool.empty) + assert(f.isDefined == true) + assert(await(f) == Seq(1, 2)) + } - "collect" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should deconstruct") { + val h = new SimpleDelayedSpoolHelper + import h._ - val f = s collect { - case x if x % 2 == 0 => x * 2 - } + assert(s match { + case fst *:: rest if fst == 1 && !rest.isDefined => true + case _ => false + }) + } - assert(f.isDefined == false) // 1 != 2 mod 0 - p() = Return(2 *:: p1) - assert(f.isDefined == true) - val s1 = Await.result(f) - assert(s1 match { - case x *:: rest if x == 4 && !rest.isDefined => true - case _ => false - }) - p1() = Return(3 *:: p2) - assert(s1 match { - case x *:: rest if x == 4 && !rest.isDefined => true - case _ => false - }) - p2() = Return(4 *:: Future.value(Spool.empty[Int])) - val s1s = s1.toSeq - assert(s1s.isDefined == true) - assert(Await.result(s1s) == Seq(4, 8)) - } + test("Simple delayed Spool should collect") { + val h = new SimpleDelayedSpoolHelper + import h._ + + val f = s collect { + case x if x % 2 == 0 => x * 2 + } + + assert(f.isDefined == false) // 1 != 2 mod 0 + p() = Return(2 *:: p1) + assert(f.isDefined == true) + val s1 = await(f) + assert(s1 match { + case x *:: rest if x == 4 && !rest.isDefined => true + case _ => false + }) + p1() = Return(3 *:: p2) + assert(s1 match { + case x *:: rest if x == 4 && !rest.isDefined => true + case _ => false + }) + p2() = Return(4 *:: Future.value(Spool.empty[Int])) + val s1s = s1.toSeq + assert(s1s.isDefined == true) + assert(await(s1s) == Seq(4, 8)) + } - "fold left" in { - val h = new SimpleDelayedSpoolHelper - import h._ + test("Simple delayed Spool should fold left") { + val h = new SimpleDelayedSpoolHelper + import h._ - val f = s.foldLeft(0) { (x, y) => - x + y - } + val f = s.foldLeft(0) { (x, y) => x + y } - assert(f.isDefined == false) - p() = Return(2 *:: p1) - assert(f.isDefined == false) - p1() = Return(Spool.empty) - assert(f.isDefined == true) - assert(Await.result(f) == 3) - } + assert(f.isDefined == false) + p() = Return(2 *:: p1) + assert(f.isDefined == false) + p1() = Return(Spool.empty) + assert(f.isDefined == true) + assert(await(f) == 3) + } - "take while" in { - val h = new SimpleDelayedSpoolHelper - import h._ - - val taken = s.takeWhile(_ < 3) - assert(taken.isEmpty == false) - val f = taken.toSeq - - assert(f.isDefined == false) - p() = Return(2 *:: p1) - assert(f.isDefined == false) - p1() = Return(3 *:: p2) - // despite the Spool having an unfulfilled tail, the takeWhile is satisfied - assert(f.isDefined == true) - assert(Await.result(f) == Seq(1, 2)) - } + test("Simple delayed Spool should take while") { + val h = new SimpleDelayedSpoolHelper + import h._ + + val taken = s.takeWhile(_ < 3) + assert(taken.isEmpty == false) + val f = taken.toSeq + + assert(f.isDefined == false) + p() = Return(2 *:: p1) + assert(f.isDefined == false) + p1() = Return(3 *:: p2) + // despite the Spool having an unfulfilled tail, the takeWhile is satisfied + assert(f.isDefined == true) + assert(await(f) == Seq(1, 2)) } - // These set of tests assert that any method called on Spool consumes only the + // These next few tests assert that any method called on Spool consumes only the // head, and doesn't force the tail. An exception is collect, which consumes // up to the first defined value. This is because collect takes a // PartialFunction, which if not defined for head, recurses on the tail, // forcing it. - "Lazily evaluated Spool" should { - "be constructed lazily" in { - applyLazily(Future.value _) - } - - "collect lazily" in { - applyLazily { spool => - spool.collect { - case x if x % 2 == 0 => x - } - } - } + test("Lazily evaluated Spool should be constructed lazily") { + applyLazily(Future.value _) + } - "map lazily" in { - applyLazily { spool => - Future.value(spool.map(_ + 1)) + test("Lazily evaluated Spool should collect lazily") { + applyLazily { spool => + spool.collect { + case x if x % 2 == 0 => x } } + } - "mapFuture lazily" in { - applyLazily { spool => - spool.mapFuture(Future.value(_)) - } - } + test("Lazily evaluated Spool should map lazily") { + applyLazily { spool => Future.value(spool.map(_ + 1)) } + } - "flatMap lazily" in { - applyLazily { spool => - spool.flatMap { item => - Future.value((item to (item + 5)).toSpool) - } - } - } + test("Lazily evaluated Spool should mapFuture lazily") { + applyLazily { spool => spool.mapFuture(Future.value(_)) } + } - "takeWhile lazily" in { - applyLazily { spool => - Future.value { - spool.takeWhile(_ < Int.MaxValue) - } - } - } + test("Lazily evaluated Spool should flatMap lazily") { + applyLazily { spool => spool.flatMap { item => Future.value((item to (item + 5)).toSpool) } } + } - "take lazily" in { - applyLazily { spool => - Future.value { - spool.take(2) - } + test("Lazily evaluated Spool should takeWhile lazily") { + applyLazily { spool => + Future.value { + spool.takeWhile(_ < Int.MaxValue) } } + } - "zip lazily" in { - applyLazily { spool => - Future.value(spool.zip(spool).map { case (a, b) => a + b }) + test("Lazily evaluated Spool should take lazily") { + applyLazily { spool => + Future.value { + spool.take(2) } } + } - "act eagerly when forced" in { - val (spool, tailReached) = - applyLazily { spool => - Future.value(spool.map(_ + 1)) - } - Await.ready { spool.map(_.force) } - assert(tailReached.isDefined) - } - - /** - * Confirms that the given operation does not consume an entire Spool, and then - * returns the resulting Spool and tail check for further validation. - */ - def applyLazily(f: Spool[Int] => Future[Spool[Int]]): (Future[Spool[Int]], Future[Spool[Int]]) = { - val tail = new Promise[Spool[Int]] - - // A spool where only the head is valid. - def spool: Spool[Int] = 0 *:: { tail.setException(new Exception); tail } + test("Lazily evaluated Spool should zip lazily") { + applyLazily { spool => Future.value(spool.zip(spool).map { case (a, b) => a + b }) } + } - // create, apply, poll - val s = f(spool) - assert(!tail.isDefined) - (s, tail) - } + test("Lazily evaluated Spool should act eagerly when forced") { + val (spool, tailReached) = + applyLazily { spool => Future.value(spool.map(_ + 1)) } + Await.ready { spool.map(_.force) } + assert(tailReached.isDefined) } // Note: ++ is different from the other methods because it doesn't force its // argument. - "lazy ++" in { + test("lazy ++") { val nil = Future.value(Spool.empty[Unit]) val a = () *:: nil val p = new Promise[Unit] @@ -509,19 +478,19 @@ class SpoolTest extends WordSpec with GeneratorDrivenPropertyChecks { val c = a ++ f assert(!p.isDefined) - Await.result(c.flatMap(_.tail).select(b.tail)) + await(c.flatMap(_.tail).select(b.tail)) assert(p.isDefined) } - "Spool.merge should merge" in { + test("Spool.merge should merge") { forAll(Arbitrary.arbitrary[List[List[Int]]]) { ss => val all = Spool.merge(ss.map(s => Future.value(s.toSpool))) val interleaved = ss.flatMap(_.zipWithIndex).sortBy(_._2).map(_._1) - assert(Await.result(all.flatMap(_.toSeq)) == interleaved) + assert(await(all.flatMap(_.toSeq)) == interleaved) } } - "Spool.merge should merge round robin" in { + test("Spool.merge should merge round robin") { val spools: Seq[Future[Spool[String]]] = Seq( "a" *:: Future.value("b" *:: Future.value("c" *:: Future.value(Spool.empty[String]))), "1" *:: Future.value("2" *:: Future.value("3" *:: Future.value(Spool.empty[String]))), @@ -529,21 +498,21 @@ class SpoolTest extends WordSpec with GeneratorDrivenPropertyChecks { "foo" *:: Future.value("bar" *:: Future.value("baz" *:: Future.value(Spool.empty[String]))) ).map(Future.value) assert( - Await.result(Spool.merge(spools).flatMap(_.toSeq), 5.seconds) == + await(Spool.merge(spools).flatMap(_.toSeq)) == Seq("a", "1", "foo", "b", "2", "bar", "c", "3", "baz") ) } - "Spool.distinctBy should distinct" in { + test("Spool.distinctBy should distinct") { forAll { (s: List[Char]) => val d = s.toSpool.distinctBy(x => x).toSeq - assert(Await.result(d).toSet == s.toSet) + assert(await(d).toSet == s.toSet) } } - "Spool.distinctBy should distinct by in order" in { + test("Spool.distinctBy should distinct by in order") { val spool: Spool[String] = "ac" *:: Future.value("bbe" *:: Future.value("ab" *:: Future.value(Spool.empty[String]))) - assert(Await.result(spool.distinctBy(_.length).toSeq, 5.seconds) == Seq("ac", "bbe")) + assert(await(spool.distinctBy(_.length).toSeq) == Seq("ac", "bbe")) } } diff --git a/util-core/src/test/scala/com/twitter/concurrent/TxTest.scala b/util-core/src/test/scala/com/twitter/concurrent/TxTest.scala index 051e1e7e5f..9b872a8475 100644 --- a/util-core/src/test/scala/com/twitter/concurrent/TxTest.scala +++ b/util-core/src/test/scala/com/twitter/concurrent/TxTest.scala @@ -1,83 +1,77 @@ package com.twitter.concurrent -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - import com.twitter.util.Return +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class TxTest extends WordSpec { - "Tx.twoParty" should { - "commit when everything goes dandy" in { - val (stx, rtx) = Tx.twoParty(123) - val sf = stx.ack() - assert(sf.poll == None) - val rf = rtx.ack() - assert(sf.poll == Some(Return(Tx.Commit(())))) - assert(rf.poll == Some(Return(Tx.Commit(123)))) - } - - "abort when receiver nacks" in { - val (stx, rtx) = Tx.twoParty(123) - val sf = stx.ack() - assert(sf.poll == None) - rtx.nack() - assert(sf.poll == Some(Return(Tx.Abort))) - } +class TxTest extends AnyFunSuite { + test("Tx.twoParty should commit when everything goes dandy") { + val (stx, rtx) = Tx.twoParty(123) + val sf = stx.ack() + assert(sf.poll == None) + val rf = rtx.ack() + assert(sf.poll == Some(Return(Tx.Commit(())))) + assert(rf.poll == Some(Return(Tx.Commit(123)))) + } - "abort when sender nacks" in { - val (stx, rtx) = Tx.twoParty(123) - val rf = rtx.ack() - assert(rf.poll == None) - stx.nack() - assert(rf.poll == Some(Return(Tx.Abort))) - } + test("Tx.twoParty should abort when receiver nacks") { + val (stx, rtx) = Tx.twoParty(123) + val sf = stx.ack() + assert(sf.poll == None) + rtx.nack() + assert(sf.poll == Some(Return(Tx.Abort))) + } - "complain on ack ack" in { - val (stx, rtx) = Tx.twoParty(123) - rtx.ack() + test("Tx.twoParty should abort when sender nacks") { + val (stx, rtx) = Tx.twoParty(123) + val rf = rtx.ack() + assert(rf.poll == None) + stx.nack() + assert(rf.poll == Some(Return(Tx.Abort))) + } - assert(intercept[Exception] { - rtx.ack() - } == Tx.AlreadyAckd) - } + test("Tx.twoParty should complain on ack ack") { + val (stx, rtx) = Tx.twoParty(123) + rtx.ack() - "complain on ack nack" in { - val (stx, rtx) = Tx.twoParty(123) + assert(intercept[Exception] { rtx.ack() + } == Tx.AlreadyAckd) + } - assert(intercept[Exception] { - rtx.nack() - } == Tx.AlreadyAckd) - } + test("Tx.twoParty should complain on ack nack") { + val (stx, rtx) = Tx.twoParty(123) + rtx.ack() - "complain on nack ack" in { - val (stx, rtx) = Tx.twoParty(123) + assert(intercept[Exception] { rtx.nack() + } == Tx.AlreadyAckd) + } + + test("Tx.twoParty should complain on nack ack") { + val (stx, rtx) = Tx.twoParty(123) + rtx.nack() + + assert(intercept[Exception] { + rtx.ack() + } == Tx.AlreadyNackd) + } - assert(intercept[Exception] { - rtx.ack() - } == Tx.AlreadyNackd) - } + test("Tx.twoParty should complain on nack nack") { + val (stx, rtx) = Tx.twoParty(123) + rtx.nack() - "complain on nack nack" in { - val (stx, rtx) = Tx.twoParty(123) + assert(intercept[Exception] { rtx.nack() + } == Tx.AlreadyNackd) + } - assert(intercept[Exception] { - rtx.nack() - } == Tx.AlreadyNackd) - } + test("Tx.twoParty should complain when already done") { + val (stx, rtx) = Tx.twoParty(123) + stx.ack() + rtx.ack() - "complain when already done" in { - val (stx, rtx) = Tx.twoParty(123) + assert(intercept[Exception] { stx.ack() - rtx.ack() - - assert(intercept[Exception] { - stx.ack() - } == Tx.AlreadyDone) - } + } == Tx.AlreadyDone) } } diff --git a/util-core/src/test/scala/com/twitter/conversions/BUILD b/util-core/src/test/scala/com/twitter/conversions/BUILD new file mode 100644 index 0000000000..c001c9eddb --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/BUILD @@ -0,0 +1,19 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + ], +) diff --git a/util-core/src/test/scala/com/twitter/conversions/DurationOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/DurationOpsTest.scala new file mode 100644 index 0000000000..9e9538173d --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/DurationOpsTest.scala @@ -0,0 +1,27 @@ +package com.twitter.conversions + +import com.twitter.util.Duration +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.scalacheck.ScalaCheckPropertyChecks + +class DurationOpsTest extends AnyFunSuite with ScalaCheckPropertyChecks { + test("converts Duration.Zero") { + assert(0.seconds eq Duration.Zero) + assert(0.milliseconds eq Duration.Zero) + assert(0.seconds eq 0.seconds) + } + + test("converts nonzero durations") { + assert(1.seconds == Duration.fromSeconds(1)) + assert(123.milliseconds == Duration.fromMilliseconds(123)) + assert(100L.nanoseconds == Duration.fromNanoseconds(100L)) + } + + test("conversion works for Int values") { + forAll { (intValue: Int) => + assert(intValue.seconds == intValue.toLong.seconds) + } + } + +} diff --git a/util-core/src/test/scala/com/twitter/conversions/MapOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/MapOpsTest.scala new file mode 100644 index 0000000000..6370367e7c --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/MapOpsTest.scala @@ -0,0 +1,49 @@ +package com.twitter.conversions + +import com.twitter.conversions.MapOps._ +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.SortedMap + +class MapOpsTest extends AnyFunSuite { + + test("RichMap#mapKeys") { + val mappedKeys = Map(1 -> "a") mapKeys { _.toString } + assert(mappedKeys == Map("1" -> "a")) + } + + test("RichMap#invert simple") { + val invertedMap = Map(1 -> "a").invert + assert(invertedMap == Map("a" -> Iterable(1))) + } + + test("RichMap#invert complex") { + val complexInverted = Map(1 -> "a", 2 -> "a", 3 -> "b").invert + assert(complexInverted == Map("a" -> Iterable(1, 2), "b" -> Iterable(3))) + } + + test("RichMap#invertSingleValue") { + val singleValueInverted = Map(1 -> "a", 2 -> "a", 3 -> "b").invertSingleValue + assert(singleValueInverted == Map("a" -> 2, "b" -> 3)) + } + + test("RichMap#filterValues") { + val filteredMap = Map(1 -> "a", 2 -> "a", 3 -> "b") filterValues { _ == "b" } + assert(filteredMap == Map(3 -> "b")) + } + + test("RichMap#filterNotValues") { + val filteredMap = Map(1 -> "a", 2 -> "a", 3 -> "b") filterNotValues { _ == "b" } + assert(filteredMap == Map(1 -> "a", 2 -> "a")) + } + + test("RichMap#filterNotKeys") { + val filteredMap = Map(1 -> "a", 2 -> "a", 3 -> "b") filterNotKeys { _ == 3 } + assert(filteredMap == Map(1 -> "a", 2 -> "a")) + } + + test("RichMap#toSortedMap") { + val sortedMap = Map(6 -> "f", 2 -> "b", 1 -> "a").toSortedMap + assert(sortedMap == SortedMap(1 -> "a", 2 -> "b", 6 -> "f")) + } + +} diff --git a/util-core/src/test/scala/com/twitter/conversions/OpsTogetherTest.scala b/util-core/src/test/scala/com/twitter/conversions/OpsTogetherTest.scala new file mode 100644 index 0000000000..8d689d9c08 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/OpsTogetherTest.scala @@ -0,0 +1,16 @@ +package com.twitter.conversions + +import com.twitter.util.{Duration, StorageUnit} +import org.scalatest.funsuite.AnyFunSuite + +class OpsTogetherTest extends AnyFunSuite { + test("multiple wildcard imports") { + import com.twitter.conversions.DurationOps._ + import com.twitter.conversions.PercentOps._ + import com.twitter.conversions.StorageUnitOps._ + + assert(5.percent == 0.05) + assert(1.seconds == Duration.fromSeconds(1)) + assert(1.byte == StorageUnit.fromBytes(1)) + } +} diff --git a/util-core/src/test/scala/com/twitter/conversions/PercentTest.scala b/util-core/src/test/scala/com/twitter/conversions/PercentOpsTest.scala similarity index 71% rename from util-core/src/test/scala/com/twitter/conversions/PercentTest.scala rename to util-core/src/test/scala/com/twitter/conversions/PercentOpsTest.scala index 2777e68cfa..fbce685289 100644 --- a/util-core/src/test/scala/com/twitter/conversions/PercentTest.scala +++ b/util-core/src/test/scala/com/twitter/conversions/PercentOpsTest.scala @@ -1,12 +1,12 @@ package com.twitter.conversions +import com.twitter.conversions.PercentOps._ import org.scalacheck.Arbitrary.arbitrary import org.scalacheck.Gen -import org.scalatest.FunSuite -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -class PercentTest extends FunSuite with GeneratorDrivenPropertyChecks { - import percent._ +class PercentOpsTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { private[this] val Precision = 0.0000000001 @@ -32,13 +32,11 @@ class PercentTest extends FunSuite with GeneratorDrivenPropertyChecks { } test("assorted percentages") { - forAll(arbitrary[Int]) { i => - assert(new RichDoublePercent(i).percent == i / 100.0) - } + forAll(arbitrary[Int]) { i => assert(new RichPercent(i).percent == i / 100.0) } // We're not as accurate when we get into high double-digit exponents, but that's acceptable. - forAll(Gen.choose(-10000D, 10000D)) { d => - assert(doubleEq(new RichDoublePercent(d).percent, d / 100.0)) + forAll(Gen.choose(-10000d, 10000d)) { d => + assert(doubleEq(new RichPercent(d).percent, d / 100.0)) } } diff --git a/util-core/src/test/scala/com/twitter/conversions/SeqOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/SeqOpsTest.scala new file mode 100644 index 0000000000..88d1aad871 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/SeqOpsTest.scala @@ -0,0 +1,43 @@ +package com.twitter.conversions + +import com.twitter.conversions.SeqOps._ +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.compat.immutable.LazyList + +class SeqOpsTest extends AnyFunSuite { + + test("RichSeq#createMap") { + val map = Seq("a", "and") createMap (_.size, _.toUpperCase) + assert(map == Map(1 -> "A", 3 -> "AND")) + } + + test("RichSeq#groupBySingle chooses last element in seq when key collision occurs") { + val map = Seq("a", "and", "the") groupBySingleValue { _.size } + assert(map == Map(3 -> "the", 1 -> "a")) + } + + test("RichSeq#findItemAfter") { + assert(Seq(1, 2, 3).findItemAfter(1) == Some(2)) + assert(Seq(1, 2, 3).findItemAfter(2) == Some(3)) + assert(Seq(1, 2, 3).findItemAfter(3) == None) + assert(Seq(1, 2, 3).findItemAfter(4) == None) + assert(Seq(1, 2, 3).findItemAfter(5) == None) + assert(Seq[Int]().findItemAfter(5) == None) + + assert(LazyList(1, 2, 3).findItemAfter(1) == Some(2)) + assert(LazyList[Int]().findItemAfter(5) == None) + + assert(Vector(1, 2, 3).findItemAfter(1) == Some(2)) + assert(Vector[Int]().findItemAfter(5) == None) + } + + test("RichSeq#foreachPartial") { + var numStrings = 0 + Seq("a", 1) foreachPartial { + case str: String => + numStrings += 1 + } + assert(numStrings == 1) + } + +} diff --git a/util-core/src/test/scala/com/twitter/conversions/StorageUnitOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/StorageUnitOpsTest.scala new file mode 100644 index 0000000000..587b25cf73 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/StorageUnitOpsTest.scala @@ -0,0 +1,36 @@ +package com.twitter.conversions + +import com.twitter.util.StorageUnit +import com.twitter.conversions.StorageUnitOps._ +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.scalacheck.ScalaCheckPropertyChecks + +class StorageUnitOpsTest extends AnyFunSuite with ScalaCheckPropertyChecks { + + test("converts") { + assert(StorageUnit.fromBytes(1) == 1.byte) + assert(StorageUnit.fromBytes(2) == 2.bytes) + + assert(StorageUnit.fromKilobytes(1) == 1.kilobyte) + assert(StorageUnit.fromKilobytes(3) == 3.kilobytes) + + assert(StorageUnit.fromMegabytes(1) == 1.megabyte) + assert(StorageUnit.fromMegabytes(4) == 4.megabytes) + + assert(StorageUnit.fromGigabytes(1) == 1.gigabyte) + assert(StorageUnit.fromGigabytes(5) == 5.gigabytes) + + assert(StorageUnit.fromTerabytes(1) == 1.terabyte) + assert(StorageUnit.fromTerabytes(6) == 6.terabytes) + + assert(StorageUnit.fromPetabytes(1) == 1.petabyte) + assert(StorageUnit.fromPetabytes(7) == 7.petabytes) + } + + test("conversion works for Int values") { + forAll { (intValue: Int) => + assert(intValue.bytes == intValue.toLong.bytes) + } + } + +} diff --git a/util-core/src/test/scala/com/twitter/conversions/StringOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/StringOpsTest.scala new file mode 100644 index 0000000000..a7857e1011 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/StringOpsTest.scala @@ -0,0 +1,166 @@ +/* + * Copyright 2010 Twitter, Inc. + * + * 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 + * + * https://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 com.twitter.conversions + +import com.twitter.conversions.StringOps._ +import org.scalatest.funsuite.AnyFunSuite + +class StringOpsTest extends AnyFunSuite { + + test("string#quoteC") { + assert("nothing".quoteC() == "nothing") + assert( + "name\tvalue\t\u20acb\u00fcllet?\u20ac".quoteC() == "name\\tvalue\\t\\u20acb\\xfcllet?\\u20ac" + ) + assert("she said \"hello\"".quoteC() == "she said \\\"hello\\\"") + assert("\\backslash".quoteC() == "\\\\backslash") + } + + test("string#unquoteC") { + assert("nothing".unquoteC() == "nothing") + assert( + "name\\tvalue\\t\\u20acb\\xfcllet?\\u20ac" + .unquoteC() == "name\tvalue\t\u20acb\u00fcllet?\u20ac" + ) + assert("she said \\\"hello\\\"".unquoteC() == "she said \"hello\"") + assert("\\\\backslash".unquoteC() == "\\backslash") + assert("real\\$dollar".unquoteC() == "real\\$dollar") + assert("silly\\/quote".unquoteC() == "silly/quote") + } + + test("string#hexlify") { + assert("hello".getBytes.slice(1, 4).hexlify == "656c6c") + assert("hello".getBytes.hexlify == "68656c6c6f") + } + + test("string#unhexlify") { + assert("656c6c".unhexlify().toList == "hello".getBytes.slice(1, 4).toList) + assert("68656c6c6f".unhexlify().toList == "hello".getBytes.toList) + "5".unhexlify() + assert("5".unhexlify().hexlify.toInt == 5) + } + + test("string#toCamelCase") { + assert("foo_bar".toCamelCase == "fooBar") + + assert(toCamelCase("foo_bar") == "fooBar") + assert(toPascalCase("foo_bar") == "FooBar") + + assert(toCamelCase("foo__bar") == "fooBar") + assert(toPascalCase("foo__bar") == "FooBar") + + assert(toCamelCase("foo___bar") == "fooBar") + assert(toPascalCase("foo___bar") == "FooBar") + + assert(toCamelCase("_foo_bar") == "fooBar") + assert(toPascalCase("_foo_bar") == "FooBar") + + assert(toCamelCase("foo_bar_") == "fooBar") + assert(toPascalCase("foo_bar_") == "FooBar") + + assert(toCamelCase("_foo_bar_") == "fooBar") + assert(toPascalCase("_foo_bar_") == "FooBar") + + assert(toCamelCase("a_b_c_d") == "aBCD") + assert(toPascalCase("a_b_c_d") == "ABCD") + } + + test("string#toPascalCase") { + assert("foo_bar".toPascalCase == "FooBar") + } + + test("string#toSnakeCase") { + assert("FooBar".toSnakeCase == "foo_bar") + } + + test("string#toPascalCase return an empty string if given null") { + assert(toPascalCase(null) == "") + } + + test("string#toPascalCase leave a CamelCased name untouched") { + assert(toPascalCase("GetTweets") == "GetTweets") + assert(toPascalCase("FooBar") == "FooBar") + assert(toPascalCase("HTML") == "HTML") + assert(toPascalCase("HTML5") == "HTML5") + assert(toPascalCase("Editor2TOC") == "Editor2TOC") + } + + test("string#toCamelCase Method") { + assert(toCamelCase("GetTweets") == "getTweets") + assert(toCamelCase("FooBar") == "fooBar") + assert(toCamelCase("HTML") == "hTML") + assert(toCamelCase("HTML5") == "hTML5") + assert(toCamelCase("Editor2TOC") == "editor2TOC") + } + + test("string#toPascalCase & toCamelCase Method function") { + // test both produce empty strings + toPascalCase(null).isEmpty && toCamelCase(null).isEmpty + + val name = "emperor_norton" + // test that the first letter for toCamelCase is lower-cased + assert(toCamelCase(name).toList.head.isLower) + // test that toPascalCase and toCamelCase.capitalize are the same + assert(toPascalCase(name) == toCamelCase(name).capitalize) + } + + test("string#toSnakeCase replace upper case with underscore") { + assert(toSnakeCase("MyCamelCase") == "my_camel_case") + assert(toSnakeCase("CamelCase") == "camel_case") + assert(toSnakeCase("Camel") == "camel") + assert(toSnakeCase("MyCamel12Case") == "my_camel12_case") + assert(toSnakeCase("CamelCase12") == "camel_case12") + assert(toSnakeCase("Camel12") == "camel12") + assert(toSnakeCase("Foobar") == "foobar") + } + + test("string#toSnakeCase not modify existing snake case strings") { + assert(toSnakeCase("my_snake_case") == "my_snake_case") + assert(toSnakeCase("snake") == "snake") + } + + test("string#toSnakeCase handle abbeviations") { + assert(toSnakeCase("ABCD") == "abcd") + assert(toSnakeCase("HTML") == "html") + assert(toSnakeCase("HTMLEditor") == "html_editor") + assert(toSnakeCase("EditorTOC") == "editor_toc") + assert(toSnakeCase("HTMLEditorTOC") == "html_editor_toc") + + assert(toSnakeCase("HTML5") == "html5") + assert(toSnakeCase("HTML5Editor") == "html5_editor") + assert(toSnakeCase("Editor2TOC") == "editor2_toc") + assert(toSnakeCase("HTML5Editor2TOC") == "html5_editor2_toc") + } + + test("string#toOption when nonEmpty") { + assert(toOption("foo") == Some("foo")) + } + + test("string#toOption when empty") { + assert(toOption("") == None) + assert(toOption(null) == None) + } + + test("string#getOrElse when nonEmpty") { + assert(getOrElse("foo", "purple") == "foo") + } + + test("string#getOrElse when empty") { + assert(getOrElse("", "purple") == "purple") + assert(getOrElse(null, "purple") == "purple") + } +} diff --git a/util-core/src/test/scala/com/twitter/conversions/TimeTest.scala b/util-core/src/test/scala/com/twitter/conversions/TimeTest.scala deleted file mode 100644 index 937c0e7f30..0000000000 --- a/util-core/src/test/scala/com/twitter/conversions/TimeTest.scala +++ /dev/null @@ -1,20 +0,0 @@ -package com.twitter.conversions - -import org.scalatest.FunSuite - -import com.twitter.util.Duration - -class TimeTest extends FunSuite { - import time._ - - test("converts Duration.Zero") { - assert(0.seconds eq Duration.Zero) - assert(0.milliseconds eq Duration.Zero) - assert(0.seconds eq 0.seconds) - } - - test("converts nonzero durations") { - assert(1.seconds == Duration.fromSeconds(1)) - assert(123.milliseconds == Duration.fromMilliseconds(123)) - } -} diff --git a/util-core/src/test/scala/com/twitter/conversions/TupleOpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/TupleOpsTest.scala new file mode 100644 index 0000000000..19d6dc947e --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/TupleOpsTest.scala @@ -0,0 +1,47 @@ +package com.twitter.conversions + +import com.twitter.conversions.TupleOps._ +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.SortedMap + +class TupleOpsTest extends AnyFunSuite { + + val tuples = Seq(1 -> "bob", 2 -> "sally") + + test("RichTuple#toKeys") { + assert(tuples.toKeys == Seq(1, 2)) + } + + test("RichTuple#toKeySet") { + assert(tuples.toKeySet == Set(1, 2)) + } + + test("RichTuple#toValues") { + assert(tuples.toValues == Seq("bob", "sally")) + } + + test("RichTuple#mapValues") { + assert(tuples.mapValues { _.length } == Seq(1 -> 3, 2 -> 5)) + } + + test("RichTuple#groupByKey") { + val multiTuples = Seq(1 -> "a", 1 -> "b", 2 -> "ab") + assert(multiTuples.groupByKey == Map(2 -> Seq("ab"), 1 -> Seq("a", "b"))) + } + + test("RichTuple#groupByKeyAndReduce") { + val multiTuples = Seq(1 -> 5, 1 -> 6, 2 -> 7) + assert(multiTuples.groupByKeyAndReduce(_ + _) == Map(2 -> 7, 1 -> 11)) + } + + test("RichTuple#sortByKey") { + val unorderedTuples = Seq(2 -> "sally", 1 -> "bob") + assert(unorderedTuples.sortByKey == Seq(1 -> "bob", 2 -> "sally")) + } + + test("RichTuple#toSortedMap") { + val unorderedTuples = Seq(2 -> "sally", 1 -> "bob") + assert(unorderedTuples.toSortedMap == SortedMap(1 -> "bob", 2 -> "sally")) + } + +} diff --git a/util-core/src/test/scala/com/twitter/conversions/U64OpsTest.scala b/util-core/src/test/scala/com/twitter/conversions/U64OpsTest.scala new file mode 100644 index 0000000000..f6c20fc168 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/conversions/U64OpsTest.scala @@ -0,0 +1,25 @@ +package com.twitter.conversions + +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite + +class U64OpsTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { + import com.twitter.conversions.U64Ops._ + + test("toU64HextString") { + forAll { (l: Long) => assert(l.toU64HexString == "%016x".format(l)) } + } + + test("toU64Long") { + forAll { (l: Long) => + val s = "%016x".format(l) + assert(s.toU64Long == java.lang.Long.parseUnsignedLong(s, 16)) + } + } + + test("conversion works for Int values") { + forAll { (intValue: Int) => + assert(intValue.toU64HexString == intValue.toLong.toU64HexString) + } + } +} diff --git a/util-core/src/test/scala/com/twitter/io/ActivitySourceTest.scala b/util-core/src/test/scala/com/twitter/io/ActivitySourceTest.scala new file mode 100644 index 0000000000..32199c9fa5 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/ActivitySourceTest.scala @@ -0,0 +1,32 @@ +package com.twitter.io + +import com.twitter.util.Activity +import org.scalatest.funsuite.AnyFunSuite + +class ActivitySourceTest extends AnyFunSuite { + + val ok: ActivitySource[String] = new ActivitySource[String] { + def get(varName: String) = Activity.value(varName) + } + + val failed: ActivitySource[Nothing] = new ActivitySource[Nothing] { + def get(varName: String) = Activity.exception(new Exception(varName)) + } + + test("ActivitySource.orElse") { + val a = failed.orElse(ok) + val b = ok.orElse(failed) + + assert("42" == a.get("42").sample()) + assert("42" == b.get("42").sample()) + } + +} + +object ActivitySourceTest { + + def bufToString(buf: Buf): String = buf match { + case Buf.Utf8(s) => s + } + +} diff --git a/util-core/src/test/scala/com/twitter/io/BUILD b/util-core/src/test/scala/com/twitter/io/BUILD new file mode 100644 index 0000000000..3e11414988 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/BUILD @@ -0,0 +1,24 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-core/src/test/resources", + ], +) diff --git a/util-core/src/test/scala/com/twitter/io/BufByteWriterTest.scala b/util-core/src/test/scala/com/twitter/io/BufByteWriterTest.scala index 422e44b5d9..7b835ba818 100644 --- a/util-core/src/test/scala/com/twitter/io/BufByteWriterTest.scala +++ b/util-core/src/test/scala/com/twitter/io/BufByteWriterTest.scala @@ -2,10 +2,10 @@ package com.twitter.io import java.lang.{Double => JDouble, Float => JFloat} import java.nio.charset.StandardCharsets -import org.scalatest.FunSuite -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyChecks { +final class BufByteWriterTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { import ByteWriter.OverflowException private[this] def assertIndex(bw: ByteWriter, index: Int): Unit = bw match { @@ -33,7 +33,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck } def testWriteByte(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeByte")(forAll { byte: Byte => + test(s"$name: writeByte")(forAll { (byte: Byte) => val bw = bwFactory() val buf = bw.writeByte(byte).owned() @@ -43,7 +43,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteShort(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeShort{BE,LE}")(forAll { s: Short => + test(s"$name: writeShort{BE,LE}")(forAll { (s: Short) => val be = bwFactory().writeShortBE(s) val le = bwFactory().writeShortLE(s) @@ -64,7 +64,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteMedium(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeMedium{BE,LE}")(forAll { m: Int => + test(s"$name: writeMedium{BE,LE}")(forAll { (m: Int) => val be = bwFactory().writeMediumBE(m) val le = bwFactory().writeMediumLE(m) @@ -86,7 +86,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteInt(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeInt{BE,LE}")(forAll { i: Int => + test(s"$name: writeInt{BE,LE}")(forAll { (i: Int) => val be = bwFactory().writeIntBE(i) val le = bwFactory().writeIntLE(i) @@ -109,7 +109,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteLong(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeLong{BE,LE}")(forAll { l: Long => + test(s"$name: writeLong{BE,LE}")(forAll { (l: Long) => val be = bwFactory().writeLongBE(l) val le = bwFactory().writeLongLE(l) @@ -136,7 +136,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteFloat(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeFloat{BE,LE}")(forAll { f: Float => + test(s"$name: writeFloat{BE,LE}")(forAll { (f: Float) => val be = bwFactory().writeFloatBE(f) val le = bwFactory().writeFloatLE(f) @@ -161,7 +161,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck }) def testWriteDouble(name: String, bwFactory: () => BufByteWriter, overflowOK: Boolean): Unit = - test(s"$name: writeDouble{BE,LE}")(forAll { d: Double => + test(s"$name: writeDouble{BE,LE}")(forAll { (d: Double) => val be = bwFactory().writeDoubleBE(d) val le = bwFactory().writeDoubleLE(d) @@ -195,9 +195,9 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck } test("trims unwritten bytes") { - assert(BufByteWriter.fixed(5).owned.length == 0) - assert(BufByteWriter.fixed(5).writeIntBE(1).owned.length == 4) - assert(BufByteWriter.fixed(4).writeIntBE(1).owned.length == 4) + assert(BufByteWriter.fixed(5).owned().length == 0) + assert(BufByteWriter.fixed(5).writeIntBE(1).owned().length == 4) + assert(BufByteWriter.fixed(4).writeIntBE(1).owned().length == 4) } testWriteString("fixed", size => BufByteWriter.fixed(size), overflowOK = false) @@ -209,7 +209,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck testWriteFloat("fixed", () => BufByteWriter.fixed(4), overflowOK = false) testWriteDouble("fixed", () => BufByteWriter.fixed(8), overflowOK = false) - test("fixed: writeBytes(Array[Byte])")(forAll { bytes: Array[Byte] => + test("fixed: writeBytes(Array[Byte])")(forAll { (bytes: Array[Byte]) => val bw = BufByteWriter.fixed(bytes.length) val buf = bw.writeBytes(bytes).owned() intercept[OverflowException] { bw.writeByte(0xff) } @@ -217,7 +217,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck assertIndex(bw, bytes.length) }) - test("fixed: writeBytes(Array[Byte]) 2 times")(forAll { bytes: Array[Byte] => + test("fixed: writeBytes(Array[Byte]) 2 times")(forAll { (bytes: Array[Byte]) => val bw = BufByteWriter.fixed(bytes.length * 2) val buf = bw.writeBytes(bytes).writeBytes(bytes).owned() intercept[OverflowException] { bw.writeByte(0xff) } @@ -226,7 +226,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck assertIndex(bw, bytes.length * 2) }) - test("fixed: writeBytes(Buf)")(forAll { arr: Array[Byte] => + test("fixed: writeBytes(Buf)")(forAll { (arr: Array[Byte]) => val bytes = Buf.ByteArray.Owned(arr) val bw = BufByteWriter.fixed(bytes.length) val buf = bw.writeBytes(bytes).owned() @@ -235,7 +235,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck assertIndex(bw, bytes.length) }) - test("fixed: writeBytes(Buf) 2 times")(forAll { arr: Array[Byte] => + test("fixed: writeBytes(Buf) 2 times")(forAll { (arr: Array[Byte]) => val bytes = Buf.ByteArray.Owned(arr) val bw = BufByteWriter.fixed(bytes.length * 2) val buf = bw.writeBytes(bytes).writeBytes(bytes).owned() @@ -263,20 +263,20 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck overflowOK = true ) - test("dynamic: writeBytes(Array[Byte])")(forAll { bytes: Array[Byte] => + test("dynamic: writeBytes(Array[Byte])")(forAll { (bytes: Array[Byte]) => val bw = BufByteWriter.dynamic() val buf = bw.writeBytes(bytes).owned() assert(buf == Buf.ByteArray.Owned(bytes)) }) - test("dynamic: writeBytes(Buf)")(forAll { arr: Array[Byte] => + test("dynamic: writeBytes(Buf)")(forAll { (arr: Array[Byte]) => val bytes = Buf.ByteArray.Owned(arr) val bw = BufByteWriter.dynamic() val buf = bw.writeBytes(bytes).owned() assert(buf == bytes) }) - test("dynamic: writeBytes(Array[Byte]) 3 times")(forAll { bytes: Array[Byte] => + test("dynamic: writeBytes(Array[Byte]) 3 times")(forAll { (bytes: Array[Byte]) => val bw = BufByteWriter.dynamic() val buf = bw .writeBytes(bytes) @@ -286,7 +286,7 @@ final class BufByteWriterTest extends FunSuite with GeneratorDrivenPropertyCheck assert(buf == Buf.ByteArray.Owned(bytes ++ bytes ++ bytes)) }) - test("dynamic: writeBytes(Buf) 3 times")(forAll { arr: Array[Byte] => + test("dynamic: writeBytes(Buf) 3 times")(forAll { (arr: Array[Byte]) => val bytes = Buf.ByteArray.Owned(arr) val bw = BufByteWriter.dynamic() val buf = bw diff --git a/util-core/src/test/scala/com/twitter/io/BufInputStreamTest.scala b/util-core/src/test/scala/com/twitter/io/BufInputStreamTest.scala index c18c222d2c..ffc103071f 100644 --- a/util-core/src/test/scala/com/twitter/io/BufInputStreamTest.scala +++ b/util-core/src/test/scala/com/twitter/io/BufInputStreamTest.scala @@ -1,11 +1,8 @@ package com.twitter.io -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class BufInputStreamTest extends FunSuite { +class BufInputStreamTest extends AnyFunSuite { private[this] val fileString = "Test_All_Tests\nTest_java_io_BufferedInputStream\nTest_java_io_BufferedOutputStream\nTest_ByteArrayInputStream\nTest_java_io_ByteArrayOutputStream\nTest_java_io_DataInputStream\n" private[this] val fileBuf = Buf.ByteArray.Owned(fileString.getBytes) diff --git a/util-core/src/test/scala/com/twitter/io/BufReaderTest.scala b/util-core/src/test/scala/com/twitter/io/BufReaderTest.scala index 4e5ffb0765..ccdca2b0aa 100644 --- a/util-core/src/test/scala/com/twitter/io/BufReaderTest.scala +++ b/util-core/src/test/scala/com/twitter/io/BufReaderTest.scala @@ -1,31 +1,138 @@ package com.twitter.io -import com.twitter.util.Await -import org.junit.runner.RunWith -import org.scalacheck.Prop -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.Checkers - -@RunWith(classOf[JUnitRunner]) -class BufReaderTest extends FunSuite with Checkers { - import Prop.{forAll, throws} - - test("BufReader") { - check(forAll { (bytes: String) => +import com.twitter.conversions.DurationOps._ +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.util.{Await, Future} +import org.scalacheck.{Arbitrary, Gen} +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import scala.annotation.tailrec +import org.scalatest.funsuite.AnyFunSuite + +class BufReaderTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { + + private def await[A](f: Future[A]): A = Await.result(f, 5.seconds) + + test("BufReader - apply(buf)") { + forAll { (bytes: String) => val buf = Buf.Utf8(bytes) val r = BufReader(buf) - Await.result(Reader.readAll(r)) == buf - }) - } - - test("BufReader - discard") { - check(forAll { (bytes: String, n: Int) => - val r = BufReader(Buf.Utf8(bytes)) - r.discard() - n < 0 || - bytes.length == 0 || - throws(classOf[Reader.ReaderDiscarded])({ Await.result(r.read(n)) }) - }) + assert(await(BufReader.readAll(r)) == buf) + } + } + + test("BufReader - apply(buf, chunkSize)") { + val stringAndChunk = for { + s <- Gen.alphaStr + i <- Gen.posNum[Int].suchThat(_ <= s.length) + } yield (s, i) + + forAll(stringAndChunk) { + case (bytes, chunkSize) => + val buf = Buf.Utf8(bytes) + val r = BufReader(buf, chunkSize) + assert(await(BufReader.readAll(r)) == buf) + } + } + + test("BufReader - readAll") { + forAll { (bytes: String) => + val r = Reader.fromBuf(Buf.Utf8(bytes)) + assert(await(BufReader.readAll(r)) == Buf.Utf8(bytes)) + } + } + + test("BufReader.chunked") { + val stringAndChunk = for { + s <- Gen.alphaStr + i <- Gen.posNum[Int].suchThat(_ <= s.length) + } yield (s, i) + + forAll(stringAndChunk) { + case (s, i) => + val r = BufReader.chunked(Reader.fromBuf(Buf.Utf8(s), 32), i) + assert(await(BufReader.readAll(r)) == Buf.Utf8(s)) + } + } + + test("BufReader.chunked by the chunkSize") { + val stringAndChunk = for { + s <- Gen.alphaStr + i <- Gen.posNum[Int].suchThat(_ <= s.length) + } yield (s, i) + + forAll(stringAndChunk) { + case (s, i) => + val r = BufReader.chunked(Reader.fromBuf(Buf.Utf8(s), 32), i) + + @tailrec + def readLoop(): Unit = await(r.read()) match { + case Some(b) => + assert(b.length <= i) + readLoop() + case None => () + } + + readLoop() + } + } + + test("BufReader.framed reads framed data") { + val getByteArrays: Gen[Seq[Buf]] = Gen.listOf( + for { + // limit arrays to a few kilobytes, otherwise we may generate a very large amount of data + numBytes <- Gen.choose(0.bytes.inBytes, 2.kilobytes.inBytes) + bytes <- Gen.containerOfN[Array, Byte](numBytes.toInt, Arbitrary.arbitrary[Byte]) + } yield Buf.ByteArray.Owned(bytes) + ) + + forAll(getByteArrays) { (buffers: Seq[Buf]) => + val buffersWithLength = buffers.map(buf => Buf.U32BE(buf.length).concat(buf)) + + val r = BufReader.framed(BufReader(Buf(buffersWithLength)), new BufReaderTest.U32BEFramer()) + + // read all of the frames + buffers.foreach { buf => assert(await(r.read()).contains(buf)) } + + // make sure the reader signals EOF + assert(await(r.read()).isEmpty) + } + } + + test("BufReader.framed reads empty frames") { + val r = BufReader.framed(BufReader(Buf.U32BE(0)), new BufReaderTest.U32BEFramer()) + assert(await(r.read()).contains(Buf.Empty)) + assert(await(r.read()).isEmpty) + } + + test("Framed works on non-list collections") { + val r = BufReader.framed(BufReader(Buf.U32BE(0)), i => Vector(i)) + assert(await(r.read()).isDefined) + } +} + +object BufReaderTest { + + /** + * Used to test BufReader.framed, extract fields in terms of + * frames, signified by a 32-bit BE value preceding + * each frame. + */ + private class U32BEFramer() extends (Buf => Seq[Buf]) { + var state: Buf = Buf.Empty + + @tailrec + private def loop(acc: Seq[Buf], buf: Buf): Seq[Buf] = { + buf match { + case Buf.U32BE(l, d) if d.length >= l => + loop(acc :+ d.slice(0, l), d.slice(l, d.length)) + case _ => + state = buf + acc + } + } + + def apply(buf: Buf): Seq[Buf] = synchronized { + loop(Seq.empty, state concat buf) + } } } diff --git a/util-core/src/test/scala/com/twitter/io/BufTest.scala b/util-core/src/test/scala/com/twitter/io/BufTest.scala index e13b76a83b..948c3b1293 100644 --- a/util-core/src/test/scala/com/twitter/io/BufTest.scala +++ b/util-core/src/test/scala/com/twitter/io/BufTest.scala @@ -5,20 +5,22 @@ import java.nio.{ByteBuffer, CharBuffer} import java.nio.charset.{StandardCharsets => JChar} import java.util.Arrays import org.scalacheck.{Arbitrary, Gen, Prop} -import org.scalatest.FunSuite -import org.scalatest.junit.AssertionsForJUnit -import org.scalatest.mockito.MockitoSugar -import org.scalatest.prop.{Checkers, GeneratorDrivenPropertyChecks} +import org.scalatestplus.junit.AssertionsForJUnit +import org.scalatestplus.mockito.MockitoSugar +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatestplus.scalacheck.Checkers import scala.collection.mutable +import java.nio.charset.Charset +import org.scalatest.funsuite.AnyFunSuite class BufTest - extends FunSuite + extends AnyFunSuite with MockitoSugar - with GeneratorDrivenPropertyChecks + with ScalaCheckDrivenPropertyChecks with Checkers with AssertionsForJUnit { - val AllCharsets = Seq( + val AllCharsets: Seq[Charset] = Seq( JChar.ISO_8859_1, JChar.US_ASCII, JChar.UTF_8, @@ -69,14 +71,15 @@ class BufTest test("Buf.ByteArray.[Owned|Shared](data, begin, end) validates non-negative indexes") { val data = Array[Byte]() - Seq(-1 -> 1, 1 -> -1).foreach { case (begin, end) => - intercept[IllegalArgumentException] { - Buf.ByteArray.Owned(data, begin, end) - } + Seq(-1 -> 1, 1 -> -1).foreach { + case (begin, end) => + intercept[IllegalArgumentException] { + Buf.ByteArray.Owned(data, begin, end) + } - intercept[IllegalArgumentException] { - Buf.ByteArray.Shared(data, begin, end) - } + intercept[IllegalArgumentException] { + Buf.ByteArray.Shared(data, begin, end) + } } } @@ -101,6 +104,22 @@ class BufTest } } + test("Buf.equals") { + val bytes = Array[Byte](0, 1, 2) + val byteBuffer = Buf.ByteBuffer.Owned(java.nio.ByteBuffer.wrap(bytes)) + val byteArray = Buf.ByteArray(bytes.toSeq: _*) + val composite = Buf.ByteArray(0).concat(Buf.ByteArray(1)).concat(Buf.ByteArray(2)) + val bufs = Seq(byteBuffer, byteArray, composite) + + bufs.foreach { b0 => + bufs.foreach { b1 => + withClue(s"${b0.getClass.getName} and ${b1.getClass.getName}") { + assert(b0 == b1) + } + } + } + } + test("Buf.write(Array[Byte], offset) validates output array is large enough") { // we skip Empty because we cannot create an array // too small for this test. @@ -234,13 +253,8 @@ class BufTest } test("Buf.process from and until returns -1 when fully processed") { - def assertAllProcessed( - expected: Array[Byte], - buf: Buf, - from: Int, - until: Int - ): Unit = { - var seen = new mutable.ListBuffer[Byte]() + def assertAllProcessed(expected: Array[Byte], buf: Buf, from: Int, until: Int): Unit = { + val seen = new mutable.ListBuffer[Byte]() val processor = new Buf.Processor { def apply(byte: Byte): Boolean = { seen += byte @@ -252,7 +266,6 @@ class BufTest } val arr = Array.range(0, 9).map(_.toByte) - val byteArrayOffset = 3 val byteArray = Buf.ByteArray.Owned(arr, 3, 8) assertAllProcessed(Array[Byte](5, 6, 7), byteArray, 2, 5) @@ -444,7 +457,7 @@ class BufTest val b1 = Buf.ByteArray.Owned(a1) val b2 = Buf.ByteArray.Owned(a2) val b3 = Buf.ByteArray.Owned(a3) - val b4 = Buf.ByteArray.Owned(a4) + val _ = Buf.ByteArray.Owned(a4) val comp2 = b1.concat(b2) val comp3 = comp2.concat(b3) @@ -618,26 +631,6 @@ class BufTest assert(Buf.Utf8("abcabc") == Buf(Seq(abc, Buf.Empty, abc))) } - test("Buf.Composite.unapply") { - Buf.Empty match { - case Buf.Composite(_) => fail() - case _ => - } - - val abc = Buf.Utf8("abc") - abc match { - case Buf.Composite(_) => fail() - case _ => - } - - val xyz = Buf.Utf8("xyz") - abc.concat(xyz) match { - case Buf.Composite(bufs) => - assert(bufs == IndexedSeq(abc, xyz)) - case _ => fail() - } - } - test("Buf.Utf8: English") { val buf = Buf.Utf8("Hello, world!") assert(buf.length == 13) @@ -676,9 +669,7 @@ class BufTest def slice(i: Int, j: Int): Buf = throw new Exception("not implemented") def length: Int = 12 def write(output: Array[Byte], off: Int): Unit = - (off until off + length) foreach { i => - output(i) = 'a'.toByte - } + (off until off + length) foreach { i => output(i) = 'a'.toByte } def write(output: ByteBuffer): Unit = ??? def get(index: Int): Byte = ??? @@ -705,7 +696,8 @@ class BufTest } AllCharsets foreach { charset => - test("Buf.StringCoder: %s charset can encode and decode an English phrase".format(charset.name)) { + test( + "Buf.StringCoder: %s charset can encode and decode an English phrase".format(charset.name)) { val coder = new Buf.StringCoder(charset) {} val phrase = "Hello, world!" val encoded = coder(phrase) @@ -763,7 +755,7 @@ class BufTest val bb0 = java.nio.ByteBuffer.wrap(bytes) val Buf.ByteBuffer.Shared(bb1) = Buf.ByteBuffer.Owned(bb0) bb1.position(3) - assert(bb0.position == 0) + assert(bb0.position() == 0) } test("Buf.ByteBuffer.Direct.unapply is unsafe") { @@ -780,7 +772,7 @@ class BufTest val Buf.ByteBuffer(bb1) = Buf.ByteBuffer.Owned(bb0) assert(bb1.isReadOnly) bb1.position(3) - assert(bb0.position == 0) + assert(bb0.position() == 0) } test("ByteArray.coerce(ByteArray)") { @@ -880,6 +872,16 @@ class BufTest out == in && outByteBuf.getLong == in }) + check(Prop.forAll { (in1: Long, in2: Long) => + val buf = Buf.U128BE(in1, in2) + val Buf.U128BE(out1, out2, _) = buf + + val outByteBuf = java.nio.ByteBuffer.wrap(Buf.ByteArray.Owned.extract(buf)) + outByteBuf.order(java.nio.ByteOrder.BIG_ENDIAN) + + out1 == in1 && out2 == in2 && outByteBuf.getLong == in1 && outByteBuf.getLong == in2 + }) + check(Prop.forAll { (in: Long) => val buf = Buf.U64LE(in) val Buf.U64LE(out, _) = buf @@ -890,28 +892,57 @@ class BufTest out == in && outByteBuf.getLong == in }) + check(Prop.forAll { (in1: Long, in2: Long) => + val buf = Buf.U128LE(in1, in2) + val Buf.U128LE(out1, out2, _) = buf + + val outByteBuf = java.nio.ByteBuffer.wrap(Buf.ByteArray.Owned.extract(buf)) + outByteBuf.order(java.nio.ByteOrder.LITTLE_ENDIAN) + + out1 == in1 && out2 == in2 && outByteBuf.getLong == in2 && outByteBuf.getLong == in1 + }) + test("Buf num matching") { val hasMatch = Buf.Empty match { case Buf.U32BE(_, _) => true case Buf.U64BE(_, _) => true + case Buf.U128BE(_, _, _) => true case Buf.U32LE(_, _) => true case Buf.U64LE(_, _) => true + case Buf.U128LE(_, _, _) => true case _ => false } assert(!hasMatch) + // val uBigInt128 = BigInt("FFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF", 16) + val buf = Buf.Empty .concat(Buf.U32BE(Int.MaxValue)) .concat(Buf.U64BE(Long.MaxValue)) + .concat(Buf.U128BE(Long.MaxValue, Long.MaxValue)) .concat(Buf.U32LE(Int.MinValue)) .concat(Buf.U64LE(Long.MinValue)) - - val Buf.U32BE(be32, Buf.U64BE(be64, Buf.U32LE(le32, Buf.U64LE(le64, rem)))) = buf + .concat(Buf.U128LE(Long.MaxValue, Long.MaxValue)) + + val Buf.U32BE( + be32, + Buf + .U64BE( + be64, + Buf.U128BE( + be1281, + be1282, + Buf.U32LE(le32, Buf.U64LE(le64, Buf.U128LE(le1281, le1282, rem)))))) = + buf assert(be32 == Int.MaxValue) assert(be64 == Long.MaxValue) + assert(be1281 == Long.MaxValue) + assert(be1282 == Long.MaxValue) assert(le32 == Int.MinValue) assert(le64 == Long.MinValue) + assert(le1281 == Long.MaxValue) + assert(le1282 == Long.MaxValue) assert(rem == Buf.Empty) } @@ -1004,16 +1035,36 @@ class BufTest buf <- arbBuf.arbitrary } yield Buf(Seq.fill(n)(buf)) - forAll(Gen.listOf(bufGen)) { bufs: List[Buf] => - val concatLeft = bufs.foldLeft(Buf.Empty) { (l, r) => - l.concat(r) - } - val concatRight = bufs.foldRight(Buf.Empty) { (l, r) => - l.concat(r) - } + forAll(Gen.listOf(bufGen)) { (bufs: List[Buf]) => + val concatLeft = bufs.foldLeft(Buf.Empty) { (l, r) => l.concat(r) } + val concatRight = bufs.foldRight(Buf.Empty) { (l, r) => l.concat(r) } val constructor = Buf(bufs) assert(constructor == concatLeft) assert(constructor == concatRight) } } + + test("slowFromHexString handles empty string") { + val b = Buf.slowFromHexString("") + assert(b == Buf.ByteArray.Owned(new Array[Byte](0))) + } + + test("slowFromHexString throws for (some) invalid strings") { + // we don't guarantee all invalid strings throw but we do cover some of the basic ones + intercept[IllegalArgumentException] { + Buf.slowFromHexString("123") + } + + intercept[NumberFormatException] { + Buf.slowFromHexString("10ZZ") + } + } + + test("slowHexString round trips") { + forAll(arbBuf.arbitrary) { buf => + val asHex = Buf.slowHexString(buf) + val backToBuf = Buf.slowFromHexString(asHex) + assert(backToBuf == buf) + } + } } diff --git a/util-core/src/test/scala/com/twitter/io/ByteReaderTest.scala b/util-core/src/test/scala/com/twitter/io/ByteReaderTest.scala index 2105767c29..9b0618fcc2 100644 --- a/util-core/src/test/scala/com/twitter/io/ByteReaderTest.scala +++ b/util-core/src/test/scala/com/twitter/io/ByteReaderTest.scala @@ -2,19 +2,16 @@ package com.twitter.io import java.lang.{Double => JDouble, Float => JFloat} import java.nio.charset.StandardCharsets -import org.junit.runner.RunWith import org.scalacheck.Gen -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { +class ByteReaderTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { import ByteReader._ def readerWith(bytes: Byte*): ByteReader = ByteReader(Buf.ByteArray.Owned(bytes.toArray)) - def maskMedium(i: Int) = i & 0x00ffffff + def maskMedium(i: Int): Int = i & 0x00ffffff test("readBytes increments position by the number of bytes returned") { val br = readerWith() @@ -24,27 +21,25 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { test("readString underflow") { val br = readerWith() - intercept[UnderflowException] { - br.readString(1, StandardCharsets.UTF_8) - } + intercept[UnderflowException] { br.readString(1, StandardCharsets.UTF_8) } } test("readString")(forAll { (str1: String, str2: String) => val bytes1 = str1.getBytes(StandardCharsets.UTF_8) val bytes2 = str2.getBytes(StandardCharsets.UTF_8) - val br = readerWith(bytes1 ++ bytes2: _*) + val br = readerWith((bytes1 ++ bytes2).toSeq: _*) assert(br.readString(bytes1.length, StandardCharsets.UTF_8) == str1) assert(br.readString(bytes2.length, StandardCharsets.UTF_8) == str2) intercept[UnderflowException] { br.readByte() } }) - test("readByte")(forAll { byte: Byte => + test("readByte")(forAll { (byte: Byte) => val br = readerWith(byte) assert(br.readByte() == byte) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readShortBE")(forAll { s: Short => + test("readShortBE")(forAll { (s: Short) => val br = readerWith( ((s >> 8) & 0xff).toByte, ((s) & 0xff).toByte @@ -52,10 +47,10 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { // note, we need to cast here toShort so that the // MSB is interpreted as the sign bit. assert(br.readShortBE() == s) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readShortLE")(forAll { s: Short => + test("readShortLE")(forAll { (s: Short) => val br = readerWith( ((s) & 0xff).toByte, ((s >> 8) & 0xff).toByte @@ -64,30 +59,30 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { // note, we need to cast here toShort so that the // MSB is interpreted as the sign bit. assert(br.readShortLE() == s) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readUnsignedMediumBE")(forAll { m: Int => + test("readUnsignedMediumBE")(forAll { (m: Int) => val br = readerWith( ((m >> 16) & 0xff).toByte, ((m >> 8) & 0xff).toByte, ((m) & 0xff).toByte ) assert(br.readUnsignedMediumBE() == maskMedium(m)) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readUnsignedMediumLE")(forAll { m: Int => + test("readUnsignedMediumLE")(forAll { (m: Int) => val br = readerWith( ((m) & 0xff).toByte, ((m >> 8) & 0xff).toByte, ((m >> 16) & 0xff).toByte ) assert(br.readUnsignedMediumLE() == maskMedium(m)) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readIntBE")(forAll { i: Int => + test("readIntBE")(forAll { (i: Int) => val br = readerWith( ((i >> 24) & 0xff).toByte, ((i >> 16) & 0xff).toByte, @@ -95,10 +90,10 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((i) & 0xff).toByte ) assert(br.readIntBE() == i) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readIntLE")(forAll { i: Int => + test("readIntLE")(forAll { (i: Int) => val br = readerWith( ((i) & 0xff).toByte, ((i >> 8) & 0xff).toByte, @@ -106,10 +101,10 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((i >> 24) & 0xff).toByte ) assert(br.readIntLE() == i) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readLongBE")(forAll { l: Long => + test("readLongBE")(forAll { (l: Long) => val br = readerWith( ((l >> 56) & 0xff).toByte, ((l >> 48) & 0xff).toByte, @@ -121,10 +116,10 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((l) & 0xff).toByte ) assert(br.readLongBE() == l) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readLongLE")(forAll { l: Long => + test("readLongLE")(forAll { (l: Long) => val br = readerWith( ((l) & 0xff).toByte, ((l >> 8) & 0xff).toByte, @@ -136,53 +131,53 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((l >> 56) & 0xff).toByte ) assert(br.readLongLE() == l) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readUnsignedByte")(forAll { b: Byte => + test("readUnsignedByte")(forAll { (b: Byte) => val br = new ByteReaderImpl(Buf.Empty) { override def readByte() = b } assert(br.readUnsignedByte() == (b & 0xff)) }) - test("readUnsignedShortBE")(forAll { s: Short => + test("readUnsignedShortBE")(forAll { (s: Short) => val br = new ByteReaderImpl(Buf.Empty) { override def readShortBE() = s } assert(br.readUnsignedShortBE() == (s & 0xffff)) }) - test("readUnsignedShortLE")(forAll { s: Short => + test("readUnsignedShortLE")(forAll { (s: Short) => val br = new ByteReaderImpl(Buf.Empty) { override def readShortLE() = s } assert(br.readUnsignedShortLE() == (s & 0xffff)) }) - test("readMediumBE")(forAll { i: Int => + test("readMediumBE")(forAll { (i: Int) => val m = maskMedium(i) val br = new ByteReaderImpl(Buf.Empty) { override def readUnsignedMediumBE() = m } val expected = if (m > SignedMediumMax) m | 0xff000000 else m assert(br.readMediumBE() == expected) }) - test("readMediumLE")(forAll { i: Int => + test("readMediumLE")(forAll { (i: Int) => val m = maskMedium(i) val br = new ByteReaderImpl(Buf.Empty) { override def readUnsignedMediumLE() = m } val expected = if (m > SignedMediumMax) m | 0xff000000 else m assert(br.readMediumLE() == expected) }) - test("readUnsignedIntBE")(forAll { i: Int => + test("readUnsignedIntBE")(forAll { (i: Int) => val br = new ByteReaderImpl(Buf.Empty) { override def readIntBE() = i } - assert(br.readUnsignedIntBE() == (i & 0xffffffffl)) + assert(br.readUnsignedIntBE() == (i & 0xffffffffL)) }) - test("readUnsignedIntLE")(forAll { i: Int => + test("readUnsignedIntLE")(forAll { (i: Int) => val br = new ByteReaderImpl(Buf.Empty) { override def readIntLE() = i } - assert(br.readUnsignedIntLE() == (i & 0xffffffffl)) + assert(br.readUnsignedIntLE() == (i & 0xffffffffL)) }) val uInt64s: Gen[BigInt] = Gen .chooseNum(Long.MinValue, Long.MaxValue) .map(x => BigInt(x) + BigInt(2).pow(63)) - test("readUnsignedLongBE")(forAll(uInt64s) { bi: BigInt => + test("readUnsignedLongBE")(forAll(uInt64s) { (bi: BigInt) => val br = readerWith( ((bi >> 56) & 0xff).toByte, ((bi >> 48) & 0xff).toByte, @@ -194,10 +189,10 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((bi) & 0xff).toByte ) assert(br.readUnsignedLongBE() == bi) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) - test("readUnsignedLongLE")(forAll(uInt64s) { bi1: BigInt => + test("readUnsignedLongLE")(forAll(uInt64s) { (bi1: BigInt) => val bi = bi1.abs val br = readerWith( ((bi) & 0xff).toByte, @@ -210,31 +205,31 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { ((bi >> 56) & 0xff).toByte ) assert(br.readUnsignedLongLE() == bi) - val exc = intercept[UnderflowException] { br.readByte() } + intercept[UnderflowException] { br.readByte() } }) // .equals is required to handle NaN - test("readFloatBE")(forAll { i: Int => + test("readFloatBE")(forAll { (i: Int) => val br = new ByteReaderImpl(Buf.Empty) { override def readIntBE() = i } assert(br.readFloatBE().equals(JFloat.intBitsToFloat(i))) }) - test("readFloatLE")(forAll { i: Int => + test("readFloatLE")(forAll { (i: Int) => val br = new ByteReaderImpl(Buf.Empty) { override def readIntLE() = i } assert(br.readFloatLE().equals(JFloat.intBitsToFloat(i))) }) - test("readDoubleBE")(forAll { l: Long => + test("readDoubleBE")(forAll { (l: Long) => val br = new ByteReaderImpl(Buf.Empty) { override def readLongBE() = l } assert(br.readDoubleBE().equals(JDouble.longBitsToDouble(l))) }) - test("readDoubleLE")(forAll { l: Long => + test("readDoubleLE")(forAll { (l: Long) => val br = new ByteReaderImpl(Buf.Empty) { override def readLongLE() = l } assert(br.readDoubleLE().equals(JDouble.longBitsToDouble(l))) }) - test("readBytes")(forAll { bytes: Array[Byte] => + test("readBytes")(forAll { (bytes: Array[Byte]) => val buf = Buf.ByteArray.Owned(bytes ++ bytes) val br = ByteReader(buf) intercept[IllegalArgumentException] { br.readBytes(-1) } @@ -243,7 +238,7 @@ class ByteReaderTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(br.readBytes(1) == Buf.Empty) }) - test("readAll")(forAll { bytes: Array[Byte] => + test("readAll")(forAll { (bytes: Array[Byte]) => val buf = Buf.ByteArray.Owned(bytes ++ bytes) val br = ByteReader(buf) assert(br.readAll() == Buf.ByteArray.Owned(bytes ++ bytes)) diff --git a/util-core/src/test/scala/com/twitter/io/CachingActivitySourceTest.scala b/util-core/src/test/scala/com/twitter/io/CachingActivitySourceTest.scala new file mode 100644 index 0000000000..226070cd63 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/CachingActivitySourceTest.scala @@ -0,0 +1,18 @@ +package com.twitter.io + +import com.twitter.util.Activity +import scala.util.Random +import org.scalatest.funsuite.AnyFunSuite + +class CachingActivitySourceTest extends AnyFunSuite { + + test("CachingActivitySource") { + val cache = new CachingActivitySource[String](new ActivitySource[String] { + def get(varName: String) = Activity.value(Random.alphanumeric.take(10).mkString) + }) + + val a = cache.get("a") + assert(a == cache.get("a")) + } + +} diff --git a/util-core/src/test/scala/com/twitter/io/ClassLoaderActivitySourceTest.scala b/util-core/src/test/scala/com/twitter/io/ClassLoaderActivitySourceTest.scala new file mode 100644 index 0000000000..e59c1d1687 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/ClassLoaderActivitySourceTest.scala @@ -0,0 +1,21 @@ +package com.twitter.io + +import com.twitter.util.FuturePool +import java.io.ByteArrayInputStream +import org.scalatest.funsuite.AnyFunSuite + +class ClassLoaderActivitySourceTest extends AnyFunSuite { + + test("ClassLoaderActivitySource") { + val classLoader = new ClassLoader() { + override def getResourceAsStream(name: String) = + new ByteArrayInputStream(name.getBytes("UTF-8")) + } + + val loader = new ClassLoaderActivitySource(classLoader, FuturePool.immediatePool) + val bufAct = loader.get("bar baz") + + assert("bar baz" == ActivitySourceTest.bufToString(bufAct.sample())) + } + +} diff --git a/util-core/src/test/scala/com/twitter/io/ClasspathResourceTest.scala b/util-core/src/test/scala/com/twitter/io/ClasspathResourceTest.scala new file mode 100644 index 0000000000..4099eead82 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/ClasspathResourceTest.scala @@ -0,0 +1,116 @@ +package com.twitter.io + +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers + +class ClasspathResourceTest extends AnyFunSuite with Matchers { + + test("ClasspathResource#load absolute path") { + ClasspathResource.load("/test.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load path interpreted as absolute") { + ClasspathResource.load("test.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load absolute path multiple directories 1") { + // loads from /com/twitter/io/resource-test-file.txt + ClasspathResource.load("/com/twitter/io/resource-test-file.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load absolute path multiple directories 2") { + // loads from /foo/bar/test-file.txt + ClasspathResource.load("/foo/bar/test-file.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load path interpreted as absolute multiple directories 1") { + // loads from /com/twitter/io/resource-test-file.txt + ClasspathResource.load("com/twitter/io/resource-test-file.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load path interpreted as absolute multiple directories 2") { + // loads from /foo/bar/test-file.txt + ClasspathResource.load("foo/bar/test-file.txt") match { + case Some(inputStream) => + try { + inputStream.available() > 0 should be(true) + } finally { + inputStream.close() + } + case _ => fail() + } + } + + test("ClasspathResource#load does not exist 1") { + ClasspathResource.load("/does-not-exist.txt") match { + case Some(_) => fail() + case _ => + // should not return an inputstream + } + } + + test("ClasspathResource#load does not exist 2") { + ClasspathResource.load("does-not-exist.txt") match { + case Some(_) => fail() + case _ => + // should not return an inputstream + } + } + + test("ClasspathResource#load empty file 1") { + ClasspathResource.load("/empty-file.txt") match { + case Some(_) => fail() + case _ => + // should not return an inputstream + } + } + + test("ClasspathResource#load empty file 2") { + // loads from /empty-file.txt + ClasspathResource.load("empty-file.txt") match { + case Some(_) => fail() + case _ => + // should not return an inputstream + } + } +} diff --git a/util-core/src/test/scala/com/twitter/io/FilePollingActivitySourceTest.scala b/util-core/src/test/scala/com/twitter/io/FilePollingActivitySourceTest.scala new file mode 100644 index 0000000000..7c3279625e --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/FilePollingActivitySourceTest.scala @@ -0,0 +1,63 @@ +package com.twitter.io + +import com.twitter.util.{Activity, FuturePool, MockTimer, Time} +import java.io.File +import org.scalatest.BeforeAndAfter +import scala.util.Random +import org.scalatest.funsuite.AnyFunSuite + +class FilePollingActivitySourceTest extends AnyFunSuite with BeforeAndAfter { + + def writeToTempFile(s: String): Unit = { + val printer = new java.io.PrintWriter(new File(tempFile)) + try { + printer.print(s) + } finally { + printer.close() + } + } + + def deleteTempFile(): Unit = + new File(tempFile).delete() + + val tempFile: String = Random.alphanumeric.take(5).mkString + + before { + writeToTempFile("foo bar") + } + + after { + deleteTempFile() + } + + test("FilePollingActivitySource") { + Time.withCurrentTimeFrozen { timeControl => + import com.twitter.conversions.DurationOps._ + implicit val timer = new MockTimer + val file = new FilePollingActivitySource(1.microsecond, FuturePool.immediatePool) + val buf = file.get(tempFile) + + var content: Option[String] = None + + val listen = buf.run.changes.respond { + case Activity.Ok(b) => + content = Some(ActivitySourceTest.bufToString(b)) + case _ => + content = None + } + + timeControl.advance(2.microsecond) + timer.tick() + assert(content == Some("foo bar")) + + writeToTempFile("foo baz") + + timeControl.advance(2.microsecond) + timer.tick() + assert(content == Some("foo baz")) + + listen.close() + } + } + +} diff --git a/util-core/src/test/scala/com/twitter/io/FilesTest.scala b/util-core/src/test/scala/com/twitter/io/FilesTest.scala index 5f7a62eb8d..ea36ab6ce5 100644 --- a/util-core/src/test/scala/com/twitter/io/FilesTest.scala +++ b/util-core/src/test/scala/com/twitter/io/FilesTest.scala @@ -1,18 +1,11 @@ package com.twitter.io import java.io.File +import org.scalatest.funsuite.AnyFunSuite -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -import com.twitter.util.TempFolder - -@RunWith(classOf[JUnitRunner]) -class FilesTest extends WordSpec with TempFolder { - "Files" should { - - "delete" in withTempFolder { +class FilesTest extends AnyFunSuite with TempFolder { + test("Files should delete") { + withTempFolder { val tempFolder = new File(canonicalFolderName) val file = new File(tempFolder, "file") @@ -25,9 +18,7 @@ class FilesTest extends WordSpec with TempFolder { assert(subfile.createNewFile() == true) assert(Files.delete(tempFolder) == true) - Seq(file, subfile, folder, tempFolder).foreach { x => - assert(!x.exists) - } + Seq(file, subfile, folder, tempFolder).foreach { x => assert(!x.exists) } } } diff --git a/util-core/src/test/scala/com/twitter/io/InputStreamReaderTest.scala b/util-core/src/test/scala/com/twitter/io/InputStreamReaderTest.scala index 9a25539889..b085ae166b 100644 --- a/util-core/src/test/scala/com/twitter/io/InputStreamReaderTest.scala +++ b/util-core/src/test/scala/com/twitter/io/InputStreamReaderTest.scala @@ -1,51 +1,39 @@ package com.twitter.io -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, FuturePool} import java.io.ByteArrayInputStream -import org.junit.runner.RunWith import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class InputStreamReaderTest extends FunSuite { - def arr(i: Int, j: Int) = Array.range(i, j).map(_.toByte) - def buf(i: Int, j: Int) = Buf.ByteArray.Owned(arr(i, j)) +class InputStreamReaderTest extends AnyFunSuite { + def arr(i: Int, j: Int): Array[Byte] = Array.range(i, j).map(_.toByte) + def buf(i: Int, j: Int): Buf = Buf.ByteArray.Owned(arr(i, j)) test("InputStreamReader - read empty") { val a = Array.empty[Byte] val s = new ByteArrayInputStream(a) val r = new InputStreamReader(s, 4096) - val f = r.read(10) + val f = r.read() assert(Await.result(f, 5.seconds).isEmpty) } - test("InputStreamReader - read 0 bytes") { - val a = arr(0, 25) - val s = new ByteArrayInputStream(a) - val r = new InputStreamReader(s, 4096) - - val f1 = r.read(0) - assert(Await.result(f1, 5.seconds).get.isEmpty) - } - test("InputStreamReader - read positive bytes") { val a = arr(0, 25) val s = new ByteArrayInputStream(a) - val r = new InputStreamReader(s, 4096) + val r = new InputStreamReader(s, 10) - val f1 = r.read(10) + val f1 = r.read() assert(Await.result(f1, 5.seconds) == Some(buf(0, 10))) - val f2 = r.read(10) + val f2 = r.read() assert(Await.result(f2, 5.seconds) == Some(buf(10, 20))) - val f3 = r.read(10) + val f3 = r.read() assert(Await.result(f3, 5.seconds) == Some(buf(20, 25))) - val f4 = r.read(10) + val f4 = r.read() assert(Await.result(f4).isEmpty) } @@ -54,16 +42,16 @@ class InputStreamReaderTest extends FunSuite { val s = new ByteArrayInputStream(a) val r = new InputStreamReader(s, 100) - val f1 = r.read(1000) + val f1 = r.read() assert(Await.result(f1, 5.seconds) == Some(buf(0, 100))) - val f2 = r.read(1000) + val f2 = r.read() assert(Await.result(f2, 5.seconds) == Some(buf(100, 200))) - val f3 = r.read(1000) + val f3 = r.read() assert(Await.result(f3, 5.seconds) == Some(buf(200, 250))) - val f4 = r.read(1000) + val f4 = r.read() assert(Await.result(f4, 5.seconds).isEmpty) } @@ -74,31 +62,31 @@ class InputStreamReaderTest extends FunSuite { val r1 = new InputStreamReader(s1, 100) val r2 = new InputStreamReader(s2, 500) - val f1 = Reader.readAll(r1) + val f1 = BufReader.readAll(r1) assert(Await.result(f1, 5.seconds) == buf(0, 250)) - val f2 = Reader.readAll(r1) + val f2 = BufReader.readAll(r1) assert(Await.result(f2, 5.seconds).isEmpty) - val f3 = Reader.readAll(r2) + val f3 = BufReader.readAll(r2) assert(Await.result(f3, 5.seconds) == buf(0, 250)) - val f4 = Reader.readAll(r2) + val f4 = BufReader.readAll(r2) assert(Await.result(f4, 5.seconds).isEmpty) } test("InputStreamReader - discard") { val a = arr(0, 25) val s = new ByteArrayInputStream(a) - val r = new InputStreamReader(s, 4096) + val r = new InputStreamReader(s, 10) - val f1 = r.read(10) + val f1 = r.read() assert(Await.result(f1, 5.seconds) == Some(buf(0, 10))) r.discard() - val f2 = r.read(10) - intercept[Reader.ReaderDiscarded] { Await.result(f2, 5.seconds) } + val f2 = r.read() + intercept[ReaderDiscardedException] { Await.result(f2, 5.seconds) } } test("discard closes InputStream") { @@ -120,7 +108,7 @@ class InputStreamReaderTest extends FunSuite { test("reading EOF closes the InputStream") { val in = spy(new ByteArrayInputStream(arr(0, 10))) val reader = new InputStreamReader(in, 100, FuturePool.immediatePool) - val f = Reader.readAll(reader) + val f = BufReader.readAll(reader) assert(Await.result(f, 5.seconds) == buf(0, 10)) verify(in, times(1)).close() } diff --git a/util-core/src/test/scala/com/twitter/io/IteratorReaderTest.scala b/util-core/src/test/scala/com/twitter/io/IteratorReaderTest.scala new file mode 100644 index 0000000000..2f059a27d5 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/IteratorReaderTest.scala @@ -0,0 +1,30 @@ +package com.twitter.io + +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Await, Future} +import org.scalatestplus.scalacheck.Checkers +import org.scalatest.funsuite.AnyFunSuite + +class IteratorReaderTest extends AnyFunSuite with Checkers { + + private def await[A](f: Future[A]): A = Await.result(f, 5.seconds) + + test("IteratorReader") { + check { (list: List[String]) => + val r = Reader.fromIterator(list.iterator) + await(Reader.readAllItems(r)).mkString == list.mkString + } + } + + test("IteratorReader - discard") { + check { (list: List[String]) => + val r = Reader.fromIterator(list.iterator) + r.discard() + + Await + .ready(r.read(), 5.seconds).poll.exists( + _.throwable.isInstanceOf[ReaderDiscardedException] + ) + } + } +} diff --git a/util-core/src/test/scala/com/twitter/io/PipeTest.scala b/util-core/src/test/scala/com/twitter/io/PipeTest.scala index 0e47fafc9e..373f4818e6 100644 --- a/util-core/src/test/scala/com/twitter/io/PipeTest.scala +++ b/util-core/src/test/scala/com/twitter/io/PipeTest.scala @@ -1,20 +1,43 @@ package com.twitter.io -import com.twitter.conversions.time._ -import com.twitter.io.Reader.ReaderDiscarded -import com.twitter.util.{Await, Future, Return} -import org.scalatest.{FunSuite, Matchers} +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Await, Future, MockTimer, Return, Time} +import java.io.ByteArrayOutputStream +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers -class PipeTest extends FunSuite with Matchers { +class PipeTest extends AnyFunSuite with Matchers with ScalaCheckDrivenPropertyChecks { private def await[A](f: Future[A]): A = Await.result(f, 5.seconds) private def arr(i: Int, j: Int) = Array.range(i, j).map(_.toByte) private def buf(i: Int, j: Int) = Buf.ByteArray.Owned(arr(i, j)) + // Workaround methods for dealing with Scala compiler warnings: + // + // assert(f.isDone) cannot be called directly because it results in a Scala compiler + // warning: 'possible missing interpolator: detected interpolated identifier `$conforms`' + // + // This is due to the implicit evidence required for `Future.isDone` which checks to see + // whether the Future that is attempting to have `isDone` called on it conforms to the type + // of `Future[Unit]`. This is done using `Predef.$conforms` + // https://www.scala-lang.org/api/2.12.2/scala/Predef$.html#$conforms[A]:A%3C:%3CA + // + // Passing that directly to `assert` causes problems because the `$conforms` is also seen as + // an interpolated string. We get around it by evaluating first and passing the result to + // `assert`. + private def IsSuccess(f: Future[Unit]): Boolean = + f.poll.exists(_.isReturn) + + private def assertIsSuccess(f: Future[Unit]): Unit = + assert(IsSuccess(f)) + + private def assertIsNotSuccess(f: Future[Unit]): Unit = + assert(!IsSuccess(f)) + private def assertRead(r: Reader[Buf], i: Int, j: Int): Unit = { - val n = j - i - val f = r.read(n) + val f = r.read() assertRead(f, i, j) } @@ -36,22 +59,22 @@ class PipeTest extends FunSuite with Matchers { val buf = Buf.ByteArray.Owned(Array.range(i, j).map(_.toByte)) val f = w.write(buf) assert(f.isDefined) - assert(await(f.liftToTry) == Return(())) + assert(await(f.liftToTry) == Return.Unit) } private def assertWriteEmpty(w: Writer[Buf]): Unit = { val f = w.write(Buf.Empty) assert(f.isDefined) - assert(await(f.liftToTry) == Return(())) + assert(await(f.liftToTry) == Return.Unit) } private def assertReadEofAndClosed(rw: Pipe[Buf]): Unit = { assertReadNone(rw) - assert(rw.close().isDone) + assertIsSuccess(rw.close()) } private def assertReadNone(r: Reader[Buf]): Unit = - assert(await(r.read(1)).isEmpty) + assert(await(r.read()).isEmpty) private val failedEx = new RuntimeException("ʕ •ᴥ•ʔ") @@ -66,15 +89,14 @@ class PipeTest extends FunSuite with Matchers { val rw = new Pipe[Buf] val wf = rw.write(buf(0, 6)) assert(!wf.isDefined) - assertRead(rw, 0, 3) - assert(!wf.isDefined) - assertRead(rw, 3, 6) + assertRead(rw, 0, 6) assert(wf.isDefined) + assert(await(wf.liftToTry) == Return(())) } test("Reader.readAll") { val rw = new Pipe[Buf] - val all = Reader.readAll(rw) + val all = BufReader.readAll(rw) assert(!all.isDefined) assertWrite(rw, 0, 3) assertWrite(rw, 3, 6) @@ -91,34 +113,21 @@ class PipeTest extends FunSuite with Matchers { val rw = new Pipe[Buf] val wf = rw.write(buf(0, 6)) assert(!wf.isDefined) - val rf = rw.read(6) + val rf = rw.read() assert(rf.isDefined) assert(toSeq(await(rf)) == Seq.range(0, 6)) } - test("partial read, then short read") { - val rw = new Pipe[Buf] - val wf = rw.write(buf(0, 6)) - assert(!wf.isDefined) - val rf = rw.read(4) - assert(rf.isDefined) - assert(toSeq(await(rf)) == Seq.range(0, 4)) - - assert(!wf.isDefined) - val rf2 = rw.read(4) - assert(rf2.isDefined) - assert(toSeq(await(rf2)) == Seq.range(4, 6)) - - assert(wf.isDefined) - assert(await(wf.liftToTry) == Return(())) - } - test("fail while reading") { val rw = new Pipe[Buf] - val rf = rw.read(6) + var closed = false + rw.onClose.ensure { closed = true } + val rf = rw.read() assert(!rf.isDefined) + assert(!closed) val exc = new Exception rw.fail(exc) + assert(closed) assert(rf.isDefined) val exc1 = intercept[Exception] { await(rf) } assert(exc eq exc1) @@ -127,53 +136,59 @@ class PipeTest extends FunSuite with Matchers { test("fail before reading") { val rw = new Pipe[Buf] rw.fail(new Exception) - val rf = rw.read(10) + val rf = rw.read() assert(rf.isDefined) intercept[Exception] { await(rf) } } test("discard") { val rw = new Pipe[Buf] + var closed = false + rw.onClose.ensure { closed = true } rw.discard() - val rf = rw.read(10) + val rf = rw.read() assert(rf.isDefined) - intercept[Reader.ReaderDiscarded] { await(rf) } + assert(closed) + intercept[ReaderDiscardedException] { await(rf) } } test("close") { val rw = new Pipe[Buf] + var closed = false + rw.onClose.ensure { closed = true } val wf = rw.write(buf(0, 6)) before rw.close() assert(!wf.isDefined) - assert(await(rw.read(6)) == Some(buf(0, 6))) + assert(!closed) + assert(await(rw.read()).contains(buf(0, 6))) assert(!wf.isDefined) assertReadEofAndClosed(rw) + assert(closed) } - test("write then reads then close") { + test("write then reads then close") { testWriteReadClose() } + def testWriteReadClose() = { val rw = new Pipe[Buf] val wf = rw.write(buf(0, 6)) - assert(!wf.isDone) - assertRead(rw, 0, 3) - assert(!wf.isDone) - assertRead(rw, 3, 6) - assert(wf.isDone) - - assert(!rw.close().isDone) + assertIsNotSuccess(wf) + assertRead(rw, 0, 6) + assertIsSuccess(wf) + assertIsNotSuccess(rw.close()) assertReadEofAndClosed(rw) } - test("read then write then close") { + test("read then write then close") { readWriteClose() } + def readWriteClose() = { val rw = new Pipe[Buf] - val rf = rw.read(6) + val rf = rw.read() assert(!rf.isDefined) val wf = rw.write(buf(0, 6)) - assert(wf.isDone) + assertIsSuccess(wf) assertRead(rf, 0, 6) - assert(!rw.close().isDone) + assertIsNotSuccess(rw.close()) assertReadEofAndClosed(rw) } @@ -183,59 +198,30 @@ class PipeTest extends FunSuite with Matchers { assertFailedEx(rw.write(buf(0, 6))) val cf = rw.close() - assert(!cf.isDone) + assertIsNotSuccess(cf) - assertFailedEx(rw.read(1)) + assertFailedEx(rw.read()) assertFailedEx(cf) } test("write after close") { val rw = new Pipe[Buf] val cf = rw.close() - assert(!cf.isDone) + assertIsNotSuccess(cf) assertReadEofAndClosed(rw) - assert(cf.isDone) + assertIsSuccess(cf) intercept[IllegalStateException] { await(rw.write(buf(0, 1))) } } - test("write smaller buf than read is waiting for") { - val rw = new Pipe[Buf] - val rf = rw.read(6) - assert(!rf.isDefined) - - val wf = rw.write(buf(0, 5)) - assert(wf.isDone) - assertRead(rf, 0, 5) - - assert(!rw.read(1).isDefined) // nothing pending - rw.close() - assertReadEofAndClosed(rw) - } - - test("write larger buf than read is waiting for") { - val rw = new Pipe[Buf] - val rf = rw.read(3) - assert(!rf.isDefined) - - val wf = rw.write(buf(0, 6)) - assert(!wf.isDone) - assertRead(rf, 0, 3) - assert(!wf.isDone) - - assertRead(rw.read(5), 3, 6) // read the rest - assert(wf.isDone) - - assert(!rw.read(1).isDefined) // nothing pending to read - assert(rw.close().isDone) - } - test("write while write pending") { val rw = new Pipe[Buf] + var closed = false + rw.onClose.ensure { closed = true } val wf = rw.write(buf(0, 1)) - assert(!wf.isDone) + assertIsNotSuccess(wf) intercept[IllegalStateException] { await(rw.write(buf(0, 1))) @@ -243,50 +229,61 @@ class PipeTest extends FunSuite with Matchers { // the extraneous write should not mess with the 1st one. assertRead(rw, 0, 1) + assert(!closed) } test("read after fail") { val rw = new Pipe[Buf] rw.fail(failedEx) - assertFailedEx(rw.read(1)) + assertFailedEx(rw.read()) } - test("read after close with no pending reads") { + def readAfterCloseNoPendingReads() = { val rw = new Pipe[Buf] - assert(!rw.close().isDone) + assertIsNotSuccess(rw.close()) assertReadEofAndClosed(rw) } + test("read after close with no pending reads") { + readAfterCloseNoPendingReads() + } - test("read after close with pending data") { + def readAfterClosePendingData = { val rw = new Pipe[Buf] val wf = rw.write(buf(0, 1)) - assert(!wf.isDone) + assertIsNotSuccess(wf) // close before the write is satisfied wipes the pending write - assert(!rw.close().isDone) + assertIsNotSuccess(rw.close()) intercept[IllegalStateException] { await(wf) } - assertReadEofAndClosed(rw) + assertReadNone(rw) + intercept[IllegalStateException] { + await(rw.onClose) + } } + test("read after close with pending data") { readAfterClosePendingData } test("read while reading") { val rw = new Pipe[Buf] - val rf = rw.read(1) + var closed = false + rw.onClose.ensure { closed = true } + val rf = rw.read() intercept[IllegalStateException] { - await(rw.read(1)) + await(rw.read()) } assert(!rf.isDefined) + assert(!closed) } test("discard with pending read") { val rw = new Pipe[Buf] - val rf = rw.read(1) + val rf = rw.read() rw.discard() - intercept[ReaderDiscarded] { + intercept[ReaderDiscardedException] { await(rf) } } @@ -297,7 +294,7 @@ class PipeTest extends FunSuite with Matchers { val wf = rw.write(buf(0, 1)) rw.discard() - intercept[ReaderDiscarded] { + intercept[ReaderDiscardedException] { await(wf) } } @@ -305,80 +302,221 @@ class PipeTest extends FunSuite with Matchers { test("close not satisfied until writes are read") { val rw = new Pipe[Buf] val cf = rw.write(buf(0, 6)).before(rw.close()) - assert(!cf.isDone) - - assertRead(rw, 0, 3) - assert(!cf.isDone) + assertIsNotSuccess(cf) - assertRead(rw, 3, 6) - assert(!cf.isDone) + assertRead(rw, 0, 6) + assertIsNotSuccess(cf) assertReadEofAndClosed(rw) } - test("close not satisfied until reads are fulfilled") { + def closeNotSatisfiedUntillAllReadsDone() = { val rw = new Pipe[Buf] - val rf = rw.read(6) - val cf = rf.flatMap { _ => - rw.close() - } + val rf = rw.read() + val cf = rf.flatMap { _ => rw.close() } assert(!rf.isDefined) - assert(!cf.isDone) + assertIsNotSuccess(cf) - assert(rw.write(buf(0, 3)).isDone) + assertIsSuccess(rw.write(buf(0, 3))) assertRead(rf, 0, 3) - assert(!cf.isDone) + assertIsNotSuccess(cf) assertReadEofAndClosed(rw) } + test("close not satisfied until reads are fulfilled")(closeNotSatisfiedUntillAllReadsDone()) - test("close while read pending") { + def closeWhileReadPending = { val rw = new Pipe[Buf] - val rf = rw.read(6) + val rf = rw.read() assert(!rf.isDefined) - assert(rw.close().isDone) + assertIsSuccess(rw.close()) assert(rf.isDefined) } + test("close while read pending")(closeWhileReadPending) - test("close then close") { + def closeTwice() = { val rw = new Pipe[Buf] - assert(!rw.close().isDone) + assertIsNotSuccess(rw.close()) assertReadEofAndClosed(rw) - assert(rw.close().isDone) + assertIsSuccess(rw.close()) assertReadEofAndClosed(rw) } + test("close then close")(closeTwice()) test("close after fail") { val rw = new Pipe[Buf] rw.fail(failedEx) val cf = rw.close() - assert(!cf.isDone) + assertIsNotSuccess(cf) - assertFailedEx(rw.read(1)) + assertFailedEx(rw.read()) assertFailedEx(cf) } test("close before fail") { - val rw = new Pipe[Buf] - val cf = rw.close() - assert(!cf.isDone) + val timer = new MockTimer() + Time.withCurrentTimeFrozen { ctrl => + val rw = new Pipe[Buf](timer) + val cf = rw.close(1.second) + assertIsNotSuccess(cf) - rw.fail(failedEx) - assert(!cf.isDone) + ctrl.advance(1.second) + timer.tick() - assertFailedEx(rw.read(1)) - assertFailedEx(cf) + rw.fail(failedEx) + + assertFailedEx(rw.read()) + } + } + + test("close before fail within deadline") { + val timer = new MockTimer() + Time.withCurrentTimeFrozen { _ => + val rw = new Pipe[Buf](timer) + val cf = rw.close(1.second) + assertIsNotSuccess(cf) + + rw.fail(failedEx) + assertIsNotSuccess(cf) + + assertFailedEx(rw.read()) + assertFailedEx(cf) + } } test("close while write pending") { val rw = new Pipe[Buf] val wf = rw.write(buf(0, 1)) - assert(!wf.isDone) + assertIsNotSuccess(wf) val cf = rw.close() - assert(!cf.isDone) + assertIsNotSuccess(cf) intercept[IllegalStateException] { await(wf) } - assertReadEofAndClosed(rw) + assertReadNone(rw) + intercept[IllegalStateException] { + await(rw.onClose) + } + } + + test("close respects deadline") { + val mockTimer = new MockTimer() + Time.withCurrentTimeFrozen { timeCtrl => + val rw = new Pipe[Buf](mockTimer) + val wf = rw.write(buf(0, 6)) + + rw.close(1.second) + + assert(!wf.isDefined) + assert(!rw.onClose.isDefined) + + timeCtrl.advance(1.second) + mockTimer.tick() + + intercept[IllegalStateException] { + await(wf) + } + + intercept[IllegalStateException] { + await(rw.onClose) + } + assertReadNone(rw) + } + } + + test("read complete data before close deadline") { + val mockTimer = new MockTimer() + val rw = new Pipe[Buf](mockTimer) + Time.withCurrentTimeFrozen { timeCtrl => + val wf = rw.write(buf(0, 6)) before rw.close(1.second) + + assert(!wf.isDefined) + assertRead(rw, 0, 6) + assertReadNone(rw) + assert(wf.isDefined) + + timeCtrl.advance(1.second) + mockTimer.tick() + + assertReadEofAndClosed(rw) + } + } + + test("multiple reads read complete data before close deadline") { + val mockTimer = new MockTimer() + val buf = Buf.Utf8("foo") + Time.withCurrentTimeFrozen { timeCtrl => + val rw = new Pipe[Buf](mockTimer) + val writef = rw.write(buf) + + rw.close(1.second) + + assert(!writef.isDefined) + assert(await(BufReader.readAll(rw)) == buf) + assertReadNone(rw) + assert(writef.isDefined) + + timeCtrl.advance(1.second) + mockTimer.tick() + + assertReadEofAndClosed(rw) + } + } + + test("Pipe.copy - source and destination equality") { + forAll { (p: Array[Byte], q: Array[Byte], r: Array[Byte]) => + val rw = new Pipe[Buf] + val bos = new ByteArrayOutputStream + + val w = Writer.fromOutputStream(bos, 31) + val f = Pipe.copy(rw, w).ensure(w.close()) + val g = + rw.write(Buf.ByteArray.Owned(p)).before { + rw.write(Buf.ByteArray.Owned(q)).before { + rw.write(Buf.ByteArray.Owned(r)).before { + rw.close() + } + } + } + + await(Future.join(f, g)) + + val b = new ByteArrayOutputStream + b.write(p) + b.write(q) + b.write(r) + b.flush() + + bos.toByteArray should equal(b.toByteArray) + } } + + test("Pipe.copy - discard the source if cancelled") { + var cancelled = false + val src = new Reader[Int] { + def read(): Future[Option[Int]] = Future.value(Some(1)) + def discard(): Unit = cancelled = true + def onClose: Future[StreamTermination] = Future.never + } + val dest = new Pipe[Int] + Pipe.copy(src, dest) + assert(await(dest.read()) == Some(1)) + dest.discard() + assert(cancelled) + } + + test("Pipe.copy - discard the source if interrupted") { + var cancelled = false + val src = new Reader[Int] { + def read(): Future[Option[Int]] = Future.value(Some(1)) + def discard(): Unit = cancelled = true + def onClose: Future[StreamTermination] = Future.never + } + val dest = new Pipe[Int] + val f = Pipe.copy(src, dest) + assert(await(dest.read()) == Some(1)) + f.raise(new Exception("Freeze!")) + assert(cancelled) + } + } diff --git a/util-core/src/test/scala/com/twitter/io/ReaderSemanticsTest.scala b/util-core/src/test/scala/com/twitter/io/ReaderSemanticsTest.scala new file mode 100644 index 0000000000..211bc9d02a --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/ReaderSemanticsTest.scala @@ -0,0 +1,85 @@ +package com.twitter.io + +import com.twitter.concurrent.AsyncStream +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Future, Return, Await, Throw} +import java.io.ByteArrayInputStream +import scala.util.Random +import org.scalatest.funsuite.AnyFunSuite + +abstract class ReaderSemanticsTest[A] extends AnyFunSuite { + def createReader(): Reader[A] + + def loopAndCheck(reader: Reader[A], onClose: Future[StreamTermination]): Future[Unit] = { + reader.read().flatMap { + case Some(_) => + assert(!onClose.isDefined) + assert(!reader.onClose.isDefined) + loopAndCheck(reader, onClose) + case None => + assert(onClose.poll == Some(StreamTermination.FullyRead.Return)) + Future.Done + } + } + + protected def await[T](f: Future[T]): T = { + Await.result(f, 5.seconds) + } + + test("When you read a reader until the end, onClose will be satisfied") { + val reader = createReader() + val f = loopAndCheck(reader, reader.onClose) + await(f) + assert(f.poll == Some(Return.Unit)) + } + + test("When you discard a reader, onClose will be satisfied") { + val reader = createReader() + assert(!reader.onClose.isDefined) + reader.discard() + val onClose = reader.onClose + assert(await(onClose) == StreamTermination.Discarded) + } +} + +class BufReaderSemanticsTest extends ReaderSemanticsTest[Buf] { + def createReader(): Reader[Buf] = { + val buf = Buf.Utf8(Random.nextString(2.kilobytes.inBytes.toInt)) + BufReader(buf, 1.kilobyte.inBytes.toInt) + } +} + +class ConcatReaderSemanticsTest extends ReaderSemanticsTest[Buf] { + def createReader(): Reader[Buf] = { + val buf = Buf.Utf8(Random.nextString(2.kilobytes.inBytes.toInt)) + Reader.concat( + AsyncStream( + BufReader(buf, 1.kilobyte.inBytes.toInt), + BufReader(buf, 1.kilobyte.inBytes.toInt) + ) + ) + } + + test("When a piece is failed, the entire thing is failed") { + val left = new Pipe[Buf]() + val exn = new Exception("boom!") + left.fail(exn) + val buf = Buf.Utf8(Random.nextString(2.kilobytes.inBytes.toInt)) + val right = BufReader(buf, 1.kilobyte.inBytes.toInt) + val reader = Reader.concat(AsyncStream(left, right)) + val actual = intercept[Exception] { + await(reader.read()) + } + assert(actual == exn) + assert(reader.onClose.poll == Some(Throw(exn))) + } +} + +class InputStreamReaderSemanticsTest extends ReaderSemanticsTest[Buf] { + def createReader(): Reader[Buf] = { + val bytes = Random.nextString(2.kilobytes.inBytes.toInt).getBytes + val inputStream = new ByteArrayInputStream(bytes) + InputStreamReader(inputStream, 1.kilobyte.inBytes.toInt) + } +} diff --git a/util-core/src/test/scala/com/twitter/io/ReaderTest.scala b/util-core/src/test/scala/com/twitter/io/ReaderTest.scala index 9009291812..98b94de420 100644 --- a/util-core/src/test/scala/com/twitter/io/ReaderTest.scala +++ b/util-core/src/test/scala/com/twitter/io/ReaderTest.scala @@ -1,162 +1,293 @@ package com.twitter.io import com.twitter.concurrent.AsyncStream -import com.twitter.conversions.time._ -import com.twitter.conversions.storage._ -import com.twitter.util.{Await, Awaitable, Future, Promise} -import java.io.{ByteArrayInputStream, ByteArrayOutputStream} +import com.twitter.conversions.DurationOps._ +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.util.Await +import com.twitter.util.Awaitable +import com.twitter.util.Future +import com.twitter.util.Promise +import com.twitter.util.TimeoutException +import java.io.ByteArrayInputStream +import java.nio.charset.{StandardCharsets => JChar} import java.util.concurrent.atomic.AtomicBoolean import org.mockito.Mockito._ -import org.scalacheck.{Arbitrary, Gen} -import org.scalatest.concurrent.{Eventually, IntegrationPatience} -import org.scalatest.prop.GeneratorDrivenPropertyChecks -import org.scalatest.{FunSuite, Matchers} -import scala.annotation.tailrec +import org.scalacheck.Arbitrary +import org.scalacheck.Gen +import org.scalatest.concurrent.Eventually +import org.scalatest.concurrent.IntegrationPatience +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import scala.collection.compat.immutable.LazyList +import scala.collection.mutable.ArrayBuffer class ReaderTest - extends FunSuite - with GeneratorDrivenPropertyChecks + extends AnyFunSuite + with ScalaCheckDrivenPropertyChecks with Matchers with Eventually with IntegrationPatience { private def arr(i: Int, j: Int) = Array.range(i, j).map(_.toByte) + private def buf(i: Int, j: Int) = Buf.ByteArray.Owned(arr(i, j)) private def await[T](t: Awaitable[T]): T = Await.result(t, 5.seconds) - private def toSeq(b: Option[Buf]): Seq[Byte] = b match { - case None => fail("Expected full buffer") - case Some(buf) => - val a = new Array[Byte](buf.length) - buf.write(a, 0) - a.toSeq - } - def undefined: AsyncStream[Reader[Buf]] = throw new Exception + def undefinedReaders: List[Reader[Buf]] = List.empty + private def assertReadWhileReading(r: Reader[Buf]): Unit = { - val f = r.read(1) - intercept[IllegalStateException] { await(r.read(1)) } + val f = r.read() + intercept[IllegalStateException] { + await(r.read()) + } assert(!f.isDefined) } private def assertFailed(r: Reader[Buf], p: Promise[Option[Buf]]): Unit = { - val f = r.read(1) + val f = r.read() assert(!f.isDefined) p.setException(new Exception) - intercept[Exception] { await(f) } - intercept[Exception] { await(r.read(0)) } - intercept[Exception] { await(r.read(1)) } + intercept[Exception] { + await(f) + } + intercept[Exception] { + await(r.read()) + } + intercept[Exception] { + await(r.read()) + } } - test("Reader.copy - source and destination equality") { - forAll { (p: Array[Byte], q: Array[Byte], r: Array[Byte]) => - val rw = new Pipe[Buf] - val bos = new ByteArrayOutputStream - - val w = Writer.fromOutputStream(bos, 31) - val f = Reader.copy(rw, w).ensure(w.close()) - val g = - rw.write(Buf.ByteArray.Owned(p)).before { - rw.write(Buf.ByteArray.Owned(q)).before { - rw.write(Buf.ByteArray.Owned(r)).before { rw.close() } - } - } - - await(Future.join(f, g)) - - val b = new ByteArrayOutputStream - b.write(p) - b.write(q) - b.write(r) - b.flush() + private def writeLoop[T](from: List[T], to: Writer[T]): Future[Unit] = + from match { + case h :: t => to.write(h).flatMap(_ => writeLoop(t, to)) + case _ => to.close() + } - bos.toByteArray should equal(b.toByteArray) + test("Reader.concat collection of readers") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val concatedReaders = Reader.concat(readers) + val values = Reader.readAllItems(concatedReaders) + assert(await(values) == ss) } } - test("Reader.chunked") { - val stringAndChunk = for { - s <- Gen.alphaStr - i <- Gen.posNum[Int].suchThat(_ <= s.length) - } yield (s, i) - - forAll(stringAndChunk) { case (s, i) => - val r = Reader.chunked(Reader.fromBuf(Buf.Utf8(s)), i) - - def readLoop(): Unit = await(r.read(Int.MaxValue)) match { - case Some(b) => - assert(b.length <= i) - readLoop() - case None => () - } + test("Reader.concat from a collection won't force the stream") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val head = Reader.value("hmm") + Reader.concat(head +: readers) + assert(await(head.read()) == Some("hmm")) + } + } - readLoop() + test("Reader.concat from a stream of readers") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val concatedReaders = Reader.concat(AsyncStream.fromSeq(readers)) + val values = Reader.readAllItems(concatedReaders) + assert(await(values) == ss) } } - test("Reader.concat") { - forAll { ss: List[String] => - val readers = ss.map { s => - BufReader(Buf.Utf8(s)) - } - val buf = Reader.readAll(Reader.concat(AsyncStream.fromSeq(readers))) - assert(await(buf) == Buf.Utf8(ss.mkString)) + test("Reader.concat from a stream of readers forces the stream") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val head = Reader.value("hmm") + Reader.concat(AsyncStream.fromSeq(head +: readers)) + assert(await(head.read()) == None) } } - test("Reader.concat - discard") { + test("Reader.concat from a stream of readers - discard") { val p = new Promise[Option[Buf]] val head = new Reader[Buf] { - def read(n: Int): Future[Option[Buf]] = p - def discard(): Unit = p.setException(new Reader.ReaderDiscarded) + def read(): Future[Option[Buf]] = p + def discard(): Unit = p.setException(new ReaderDiscardedException) + def onClose: Future[StreamTermination] = Future.never } val reader = Reader.concat(head +:: undefined) reader.discard() assert(p.isDefined) } - test("Reader.concat - read while reading") { + test("Reader.concat from a stream of readers - read while reading") { val p = new Promise[Option[Buf]] val head = new Reader[Buf] { - def read(n: Int): Future[Option[Buf]] = p - def discard(): Unit = p.setException(new Reader.ReaderDiscarded) + def read(): Future[Option[Buf]] = p + def discard(): Unit = p.setException(new ReaderDiscardedException) + def onClose: Future[StreamTermination] = Future.never } val reader = Reader.concat(head +:: undefined) assertReadWhileReading(reader) } - test("Reader.concat - failed") { + test("Reader.concat from a stream of readers - failed") { val p = new Promise[Option[Buf]] val head = new Reader[Buf] { - def read(n: Int): Future[Option[Buf]] = p - def discard(): Unit = p.setException(new Reader.ReaderDiscarded) + def read(): Future[Option[Buf]] = p + def discard(): Unit = p.setException(new ReaderDiscardedException) + def onClose: Future[StreamTermination] = Future.never } val reader = Reader.concat(head +:: undefined) assertFailed(reader, p) } - test("Reader.concat - lazy tail") { + test("Reader.concat from a stream of readers - lazy tail") { val head = new Reader[Buf] { - def read(n: Int): Future[Option[Buf]] = Future.exception(new Exception) + def read(): Future[Option[Buf]] = Future.exception(new Exception) + def discard(): Unit = () + def onClose: Future[StreamTermination] = Future.never } val p = new Promise[Unit] + def tail: AsyncStream[Reader[Buf]] = { p.setDone() AsyncStream.empty } + val combined = Reader.concat(head +:: tail) - val buf = Reader.readAll(combined) - intercept[Exception] { await(buf) } + val buf = Reader.readAllItems(combined) + intercept[Exception] { + await(buf) + } assert(!p.isDefined) } + test("Reader.flatten") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val value = Reader.readAllItems(Reader.flatten(Reader.fromSeq(readers))) + assert(await(value) == ss) + } + } + + test("Reader#flatten") { + forAll { (ss: List[String]) => + val readers: List[Reader[String]] = ss.map(s => Reader.value(s)) + val buf = Reader.readAllItems((Reader.fromSeq(readers).flatten)) + assert(await(buf) == ss) + } + } + + test("Reader.flatten empty") { + val readers = Reader.flatten(Reader.fromSeq(undefinedReaders)) + assert(await(readers.read()) == None) + } + + test("Reader.flatten - discard") { + val p = new Promise[Option[Buf]] + + val head = new Reader[Buf] { + def read(): Future[Option[Buf]] = p + def discard(): Unit = p.setException(new ReaderDiscardedException) + def onClose: Future[StreamTermination] = Future.never + } + + val sourceReader = Reader.fromSeq(head +: undefinedReaders) + val reader = Reader.flatten(sourceReader) + reader.read() + reader.discard() + assert(p.isDefined) + assert(await(sourceReader.onClose) == StreamTermination.Discarded) + } + + test("Reader.flatten - failed") { + val p = new Promise[Option[Buf]] + val head = new Reader[Buf] { + def read(): Future[Option[Buf]] = p + def discard(): Unit = p.setException(new ReaderDiscardedException) + def onClose: Future[StreamTermination] = Future.never + } + val pReader = Reader.fromSeq(head +: undefinedReaders) + val reader = Reader.flatten(pReader) + assertFailed(reader, p) + assert(await(pReader.onClose) == StreamTermination.Discarded) + } + + test("Reader.flatten - onClose is satisfied when fully read") { + val reader = Reader.flatten(Reader.fromSeq(Seq(Reader.value(10), Reader.value(20)))) + assert(await(reader.read()) == Some(10)) + assert(await(reader.read()) == Some(20)) + assert(await(reader.read()) == None) + assert(await(reader.onClose) == StreamTermination.FullyRead) + } + + test("Reader.flatten - stop when encounter exceptions") { + val reader = Reader.flatten( + Reader.fromSeq( + Seq(Reader.value(10), Reader.exception(new Exception("stop")), Reader.value(20)))) + assert(await(reader.read()) == Some(10)) + val ex1 = intercept[Exception] { + await(reader.read()) + } + assert(ex1.getMessage == "stop") + val ex2 = intercept[Exception] { + await(reader.read()) + } + assert(ex2.getMessage == "stop") + val ex3 = intercept[Exception] { + await(reader.onClose) + } + assert(ex3.getMessage == "stop") + } + + test("Reader.flatten - listen for onClose update from parent") { + val exceptionMsg = "boom" + val p = new Pipe[Int] + val reader = Reader.flatten(Reader.value(p)) + reader.read() + p.fail(new Exception(exceptionMsg)) + val exception = intercept[Exception] { + await(reader.onClose) + } + assert(exception.getMessage == exceptionMsg) + } + + test( + "Reader.flatten - re-register `curReaderClosep` to listen to the next reader in the stream") { + val exceptionMsg = "boom" + + val r1 = Reader.value(1) + val p2 = new Promise + val i2 = p2.interruptible() + val r2 = Reader.fromFuture(i2) + + def rStreams: LazyList[Reader[Any]] = r1 #:: r2 #:: rStreams + + val reader = Reader.fromSeq(rStreams) + val rFlat = reader.flatten + val e = new Exception(exceptionMsg) + + assert(await(rFlat.read()) == Some(1)) + // re-register `curReaderClosep` to listen to r2 + rFlat.read() + i2.raise(e) + val exception = intercept[Exception] { + await(rFlat.onClose) + } + + assert(exception.getMessage == exceptionMsg) + } + + test("Reader.fromSeq works on infinite streams") { + def ones: LazyList[Int] = 1 #:: ones + val reader = Reader.fromSeq(ones) + assert(await(reader.read()) == Some(1)) + assert(await(reader.read()) == Some(1)) + assert(await(reader.read()) == Some(1)) + } + test("Reader.fromStream closes resources on EOF read") { val in = spy(new ByteArrayInputStream(arr(0, 10))) - val r = Reader.fromStream(in) - val f = Reader.readAll(r) + val r: Reader[Buf] = Reader.fromStream(in, 4) + val f = BufReader.readAll(r) assert(await(f) == buf(0, 10)) eventually { verify(in).close() @@ -165,100 +296,280 @@ class ReaderTest test("Reader.fromStream closes resources on discard") { val in = spy(new ByteArrayInputStream(arr(0, 10))) - val r = Reader.fromStream(in) + val r = Reader.fromStream(in, 4) r.discard() - eventually { - verify(in).close() - } + eventually { verify(in).close() } } test("Reader.fromAsyncStream completes when stream is empty") { - val as = AsyncStream(buf(1, 10)) + val as = AsyncStream(buf(1, 10)) ++ AsyncStream.exception[Buf](new Exception()) ++ AsyncStream( + buf(1, 10)) val r = Reader.fromAsyncStream(as) - val f = Reader.readAll(r) - assert(await(f) == buf(1, 10)) + assert(await(r.read()) == Some(buf(1, 10))) } test("Reader.fromAsyncStream fails on exceptional stream") { val as = AsyncStream.exception(new Exception()) val r = Reader.fromAsyncStream(as) - val f = Reader.readAll(r) - intercept[Exception] { await(f) } + val f = Reader.readAllItems(r) + intercept[Exception] { + await(f) + } } test("Reader.fromAsyncStream only evaluates tail when buffer is exhausted") { val tailEvaluated = new AtomicBoolean(false) + def tail: AsyncStream[Buf] = { tailEvaluated.set(true) AsyncStream.empty } - val as = AsyncStream.mk(buf(0, 10), tail) + + val as = AsyncStream.mk(buf(0, 10), AsyncStream.mk(buf(10, 20), tail)) val r = Reader.fromAsyncStream(as) // partially read the buffer - await(r.read(9)) + await(r.read()) assert(!tailEvaluated.get()) // read the rest of the buffer - await(r.read(2)) + await(r.read()) assert(tailEvaluated.get()) } - test("Reader.framed reads framed data") { - val getByteArrays: Gen[Seq[Buf]] = Gen.listOf( - for { - // limit arrays to a few kilobytes, otherwise we may generate a very large amount of data - numBytes <- Gen.choose(0.bytes.inBytes, 2.kilobytes.inBytes) - bytes <- Gen.containerOfN[Array, Byte](numBytes.toInt, Arbitrary.arbitrary[Byte]) - } yield Buf.ByteArray.Owned(bytes) - ) + test("Reader.toAsyncStream") { + forAll { (l: List[Byte]) => + val buf = Buf.ByteArray.Owned(l.toArray) + val as = Reader.toAsyncStream(Reader.fromBuf(buf, 1)) + + assert(await(as.toSeq()).map(b => Buf.ByteArray.Owned.extract(b).head) == l) + } + } + + test("Reader.flatMap") { + forAll { (s: String) => + val reader1 = Reader.fromBuf(Buf.Utf8(s), 8) + val reader2 = reader1.flatMap { buf => + val pipe = new Pipe[Buf] + pipe.write(buf).flatMap(_ => pipe.write(Buf.Empty)).flatMap(_ => pipe.close()) + pipe + } + + assert(Buf.decodeString(await(BufReader.readAll(reader2)), JChar.UTF_8) == s) + } + } + + test("Reader.flatMap: discard the intermediate will discard the output reader") { + val reader1 = Reader.fromBuf(Buf.Utf8("hi"), 8) + val reader2 = reader1.flatMap { buf => + val pipe = new Pipe[Buf] + pipe.write(buf).map(_ => pipe.discard()) + pipe + } + + await(reader2.read()) + intercept[ReaderDiscardedException] { + await(reader2.read()) + } + } + + test("Reader.flatMap: discard the origin reader will discard the output reader") { + val reader1 = Reader.fromBuf(Buf.Utf8("hi"), 1) + val samples = ArrayBuffer.empty[Reader[Buf]] + val reader2 = reader1.flatMap { buf => + val pipe = new Pipe[Buf] + pipe.write(buf).map(_ => samples += pipe).flatMap(_ => pipe.close()) + pipe + } + + await(reader2.read()) + reader1.discard() + + intercept[ReaderDiscardedException] { + await(reader2.read()) + } + assert(samples.size == 1) + assert(await(samples(0).onClose) == StreamTermination.FullyRead) + } + + test( + "Reader.flatMap: discard the output reader will discard " + + "the origin and intermediate") { + val reader1 = Reader.fromBuf(Buf.Utf8("hello"), 1) + val samples = ArrayBuffer.empty[Reader[Buf]] + val reader2 = reader1.flatMap { buf => + val pipe = new Pipe[Buf] + pipe.write(buf).map(_ => samples += pipe).flatMap(_ => pipe.close()) + pipe + } - forAll(getByteArrays) { buffers: Seq[Buf] => - val buffersWithLength = buffers.map(buf => Buf.U32BE(buf.length).concat(buf)) + await(reader2.read()) + await(reader2.read()) + reader2.discard() + + assert(samples.size == 2) + assert(await(samples(0).onClose) == StreamTermination.FullyRead) + assert(await(samples(1).onClose) == StreamTermination.Discarded) + + intercept[ReaderDiscardedException] { + await(reader1.read()) + } + } + + test("Reader.flatMap: propagate parent's onClose when parent reader failed with an exception") { + val ExceptionMessage = "boom" + val parentReader = new Pipe[Double] + val childReader = parentReader.flatMap(Reader.value) + parentReader.fail(new Exception(ExceptionMessage)) + val parentException = intercept[Exception](await(parentReader.onClose)) + val readerException = intercept[Exception](await(childReader.onClose)) + assert(parentException.getMessage == ExceptionMessage) + assert(readerException.getMessage == ExceptionMessage) + } + + test("Reader.flatMap: propagate parent's onClose when parent reader is discarded") { + val parentReader = Reader.empty[Int] + val childReader = parentReader.flatMap(Reader.value) + parentReader.discard() + assert(await(parentReader.onClose) == StreamTermination.Discarded) + assert(await(childReader.onClose) == StreamTermination.Discarded) + } - val r = Reader.framed(BufReader(Buf(buffersWithLength)), new ReaderTest.U32BEFramer()) + test("Reader.value") { + forAll { (a: AnyVal) => + val r = Reader.value(a) + assert(await(r.read()) == Some(a)) + assert(r.onClose.isDefined == false) + assert(await(r.read()) == None) + assert(await(r.onClose) == StreamTermination.FullyRead) + } + } - // read all of the frames - buffers.foreach { buf => - assert(await(r.read(Int.MaxValue)).contains(buf)) + test("Reader.exception") { + forAll { (ex: Exception) => + val r = Reader.exception(ex) + val exr = intercept[Exception] { + await(r.read()) + } + assert(exr == ex) + val exc = intercept[Exception] { + await(r.onClose) } + assert(exc == ex) + } + } - // make sure the reader signals EOF - assert(await(r.read(Int.MaxValue)).isEmpty) + test("Reader.fromFuture") { + // read regularly + val r1 = Reader.fromFuture(Future.value(1)) + assert(await(r1.read()) == Some(1)) + assert(await(r1.read()) == None) + + // discard + val r2 = Reader.fromFuture(Future.value(2)) + r2.discard() + intercept[ReaderDiscardedException] { + await(r2.read()) } } - test("Reader.framed reads empty frames") { - val r = Reader.framed(BufReader(Buf.U32BE(0)), new ReaderTest.U32BEFramer()) - assert(await(r.read(Int.MaxValue)).contains(Buf.Empty)) - assert(await(r.read(Int.MaxValue)).isEmpty) + test("Reader.fromFuture - unsatisfied") { + val p = Promise[Int]() + val r1 = Reader.fromFuture(p) + + val f1 = r1.read() + val f2 = f1.flatMap(_ => r1.read()) + + p.setValue(1) + assert(await(f1) == Some(1)) + // we already consumed the single item via the operation + // in `f1`, so `f2` should be empty since it's strictly + // sequenced after `f1`. + assert(await(f2) == None) } -} + test("Reader.map") { + forAll { (l: List[Int]) => + val pipe = new Pipe[Int] + writeLoop(l, pipe) + val reader2 = pipe.map(_.toString) + assert(await(Reader.readAllItems(reader2)) == l.map(_.toString)) + } + } + + test( + "Reader.empty " + + "- return Future.None and update closep to be FullyRead after first read") { + val reader = Reader.empty + assert(await(reader.read()) == None) + assert(await(reader.onClose) == StreamTermination.FullyRead) + } + + test( + "Reader.empty " + + "- return ReaderDiscardedException when reading from a discarded reader") { + val reader = Reader.empty + reader.discard() + intercept[ReaderDiscardedException] { + await(reader.read()) + } + assert(await(reader.onClose) == StreamTermination.Discarded) + } -object ReaderTest { - - /** - * Used to test Reader.framed, extract fields in terms of - * frames, signified by a 32-bit BE value preceding - * each frame. - */ - private class U32BEFramer() extends (Buf => Seq[Buf]) { - var state: Buf = Buf.Empty - - @tailrec - private def loop(acc: Seq[Buf], buf: Buf): Seq[Buf] = { - buf match { - case Buf.U32BE(l, d) if d.length >= l => - loop(acc :+ d.slice(0, l), d.slice(l, d.length)) - case _ => - state = buf - acc + test("Reader.readAllItems") { + val genBuf: Gen[Buf] = { + for { + // limit arrays to a few kilobytes, otherwise we may generate a very large amount of data + numBytes <- Gen.choose(0.bytes.inBytes, 2.kilobytes.inBytes) + bytes <- Gen.containerOfN[Array, Byte](numBytes.toInt, Arbitrary.arbitrary[Byte]) + } yield { + Buf.ByteArray.Owned(bytes) } } - def apply(buf: Buf): Seq[Buf] = synchronized { - loop(Seq.empty, state concat buf) + val genList = + Gen.listOf(Gen.oneOf(Gen.alphaLowerStr, Gen.alphaNumStr, Gen.choose(1, 100), genBuf)) + forAll(genList) { l => + val r = Reader.fromSeq(l) + assert(await(Reader.readAllItems(r)) == l) } } + + test("Reader.readAllItems - propagates interrupts if enabled") { + val p = new Promise[Int]() + val reader = Reader.fromFuture(p) + + val fut = Reader.readAllItemsInterruptible(reader) + fut.raise(new TimeoutException("")) + + assert(Await.result(reader.onClose) == StreamTermination.Discarded) + } + + test("Reader.readAllItems - ignores interrupts if disabled") { + val p = new Promise[Int]() + val reader = Reader.fromFuture(p) + + val fut = Reader.readAllItems(reader) + fut.raise(new TimeoutException("")) + + p.setValue(1) + assert(await(fut) == Seq(1)) + assert(await(reader.onClose) == StreamTermination.FullyRead) + } + + test("Reader - fails pipe with exception") { + val p = new Pipe[Int]() + val fut = Reader.readAllItemsInterruptible(p) + + fut.raise(new TimeoutException("")) + assertThrows[TimeoutException](await(p.read())) + } + + test("Reader - fails non-pipe with discard") { + val r = Reader.fromFuture(new Promise[Int]()) + val fut = Reader.readAllItemsInterruptible(r) + + fut.raise(new TimeoutException("")) + assertThrows[ReaderDiscardedException](await(r.read())) + } } diff --git a/util-core/src/test/scala/com/twitter/io/StreamIOTest.scala b/util-core/src/test/scala/com/twitter/io/StreamIOTest.scala index 36753b42d2..608c5a13dd 100644 --- a/util-core/src/test/scala/com/twitter/io/StreamIOTest.scala +++ b/util-core/src/test/scala/com/twitter/io/StreamIOTest.scala @@ -1,28 +1,22 @@ package com.twitter.io +import org.scalatest.funsuite.AnyFunSuite import scala.util.Random -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class StreamIOTest extends WordSpec { - "StreamIO.copy" should { - "copy the entire stream" in { - val buf = new Array[Byte](2048) - (new Random).nextBytes(buf) - val bis = new java.io.ByteArrayInputStream(buf) - val bos = new java.io.ByteArrayOutputStream() - StreamIO.copy(bis, bos) - assert(bos.toByteArray.toSeq == buf.toSeq) - } +class StreamIOTest extends AnyFunSuite { + test("StreamIO.copy should copy the entire stream") { + val buf = new Array[Byte](2048) + (new Random).nextBytes(buf) + val bis = new java.io.ByteArrayInputStream(buf) + val bos = new java.io.ByteArrayOutputStream() + StreamIO.copy(bis, bos) + assert(bos.toByteArray.toSeq == buf.toSeq) + } - "produce empty streams from empty streams" in { - val bis = new java.io.ByteArrayInputStream(new Array[Byte](0)) - val bos = new java.io.ByteArrayOutputStream() - StreamIO.copy(bis, bos) - assert(bos.size == 0) - } + test("StreamIO.copy should produce empty streams from empty streams") { + val bis = new java.io.ByteArrayInputStream(new Array[Byte](0)) + val bos = new java.io.ByteArrayOutputStream() + StreamIO.copy(bis, bos) + assert(bos.size == 0) } } diff --git a/util-core/src/test/scala/com/twitter/io/TempDirectoryTest.scala b/util-core/src/test/scala/com/twitter/io/TempDirectoryTest.scala index ea0d6be031..8087f8624e 100644 --- a/util-core/src/test/scala/com/twitter/io/TempDirectoryTest.scala +++ b/util-core/src/test/scala/com/twitter/io/TempDirectoryTest.scala @@ -1,24 +1,17 @@ package com.twitter.io -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class TempDirectoryTest extends WordSpec { - - "TempDirectory" should { - - "create a directory when deleteAtExit is true" in { - val dir = TempDirectory.create(true) - assert(dir.exists()) - assert(dir.isDirectory) - } +class TempDirectoryTest extends AnyFunSuite { + test("TempDirectory should create a directory when deleteAtExit is true") { + val dir = TempDirectory.create(true) + assert(dir.exists()) + assert(dir.isDirectory) + } - "create a directory when deleteAtExit is false" in { - val dir = TempDirectory.create(false) - assert(dir.exists()) - assert(dir.isDirectory) - } + test("create a directory when deleteAtExit is false") { + val dir = TempDirectory.create(false) + assert(dir.exists()) + assert(dir.isDirectory) } } diff --git a/util-core/src/test/scala/com/twitter/io/TempFileTest.scala b/util-core/src/test/scala/com/twitter/io/TempFileTest.scala index b20b8496c2..0a8d775581 100644 --- a/util-core/src/test/scala/com/twitter/io/TempFileTest.scala +++ b/util-core/src/test/scala/com/twitter/io/TempFileTest.scala @@ -1,31 +1,28 @@ package com.twitter.io -import java.io.{ByteArrayInputStream, DataInputStream} +import java.io.ByteArrayInputStream +import java.io.DataInputStream +import java.nio.charset.StandardCharsets import java.util.Arrays - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class TempFileTest extends WordSpec { - - "TempFile" should { - - "load resources" in { - val f1 = TempFile.fromResourcePath("/java/lang/String.class") - val f2 = TempFile.fromResourcePath(getClass, "/java/lang/String.class") - val f3 = TempFile.fromSystemResourcePath("java/lang/String.class") - - val c1 = Files.readBytes(f1) - val c2 = Files.readBytes(f2) - val c3 = Files.readBytes(f3) - - assert(Arrays.equals(c1, c2)) - assert(Arrays.equals(c2, c3)) - assert(new DataInputStream(new ByteArrayInputStream(c1)).readInt == 0xcafebabe) - } - +import org.scalatest.funsuite.AnyFunSuite + +class TempFileTest extends AnyFunSuite { + + test("TempFile should load resources") { + val f1 = TempFile.fromResourcePath("/java/lang/String.class") + val f2 = TempFile.fromResourcePath(getClass, "/java/lang/String.class") + val f3 = TempFile.fromSystemResourcePath("java/lang/String.class") + // Tests the basename allows basename to be less than 3 characters. + val f4 = TempFile.fromResourcePath("/1.txt") + + val c1 = Files.readBytes(f1) + val c2 = Files.readBytes(f2) + val c3 = Files.readBytes(f3) + val c4 = Files.readBytes(f4) + + assert(Arrays.equals(c1, c2)) + assert(Arrays.equals(c2, c3)) + assert(new DataInputStream(new ByteArrayInputStream(c1)).readInt == 0xcafebabe) + assert(new String(c4, StandardCharsets.UTF_8) == "Lorem ipsum dolor sit amet\n") } - } diff --git a/util-core/src/test/scala/com/twitter/io/Utf8BytesTest.scala b/util-core/src/test/scala/com/twitter/io/Utf8BytesTest.scala new file mode 100644 index 0000000000..726a4a50cb --- /dev/null +++ b/util-core/src/test/scala/com/twitter/io/Utf8BytesTest.scala @@ -0,0 +1,76 @@ +package com.twitter.io + +import org.scalatest.funsuite.AnyFunSuite +import java.nio.charset.StandardCharsets + +class Utf8BytesTest extends AnyFunSuite { + test("Utf8 created from a byte[] is immutable") { + val arr = "this is a test".getBytes(StandardCharsets.UTF_8) + val utf8 = Utf8Bytes(arr) + arr(0) = 'd' + + assert(utf8.toString == "this is a test") + } + + test("from string") { + val utf8 = Utf8Bytes("test string") + assert(utf8.toString == "test string") + } + + test("from Buf") { + val utf8 = Utf8Bytes(Buf.Utf8("test string")) + assert(utf8.toString == "test string") + } + + test("byte length and length accurate when using multi-byte characters") { + val utf8 = Utf8Bytes("\u93E1") + assert(utf8.length() == 1) + assert(utf8.asBuf.length == 3) + } + + test("simple subsequence") { + val utf8 = Utf8Bytes("test string") + val slice = utf8.subSequence(5, 11) + assert(slice == "string") + } + + test("subsequence with multi-byte characters") { + val utf8 = Utf8Bytes("\u93E1 string") + val slice = utf8.subSequence(2, 8) + assert(slice == "string") + } + + test("simple charAt") { + val utf8 = Utf8Bytes("abcdef") + assert(utf8.charAt(4) == 'e') + } + + test("charAt with multi-byte characters") { + val utf8 = Utf8Bytes("\u93E1abcdef") + assert(utf8.charAt(4) == 'd') + } + + test("asBuf round trip") { + val buf = Buf.Utf8("test string") + val utf8 = Utf8Bytes(buf) + assert(utf8.asBuf == buf) + } + + test("invalid utf8 bytes") { + val buf = Buf.ByteArray.Owned(Array(0x80.toByte)) + val utf8 = Utf8Bytes(buf) + assert(utf8.toString == "\uFFFD") + } + + test("equality and hash code") { + val a = Utf8Bytes("test string") + val b = Utf8Bytes("test string") + val c = Utf8Bytes("test string 2") + + assert(a.hashCode() == b.hashCode()) + + assert(a == a) + assert(a == b) + assert(a != c) + } +} diff --git a/util-core/src/test/scala/com/twitter/io/WriterTest.scala b/util-core/src/test/scala/com/twitter/io/WriterTest.scala index cb986aa4c5..71d216be0f 100644 --- a/util-core/src/test/scala/com/twitter/io/WriterTest.scala +++ b/util-core/src/test/scala/com/twitter/io/WriterTest.scala @@ -1,17 +1,18 @@ package com.twitter.io -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, Future, Promise} import java.io.{ByteArrayOutputStream, OutputStream} -import org.scalatest.{FunSuite, Matchers} +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers -class WriterTest extends FunSuite with Matchers { +class WriterTest extends AnyFunSuite with Matchers { private def await[A](f: Future[A]): A = Await.result(f, 5.seconds) test("fromOutputStream - close") { val w = Writer.fromOutputStream(new ByteArrayOutputStream) - w.close() + await(w.close()) intercept[IllegalStateException] { await(w.write(Buf.Empty)) } } @@ -21,6 +22,13 @@ class WriterTest extends FunSuite with Matchers { intercept[Exception] { await(w.write(Buf.Empty)) } } + test("fromOutputStream - contramap") { + val w = Writer.fromOutputStream(new ByteArrayOutputStream) + val nw = w.contramap[String](s => Buf.Utf8(s)) + await(nw.write("hi")) + await(nw.close()) + } + test("fromOutputStream - error") { val os = new OutputStream { def write(b: Int): Unit = () diff --git a/util-core/src/test/scala/com/twitter/io/exp/ActivitySourceTest.scala b/util-core/src/test/scala/com/twitter/io/exp/ActivitySourceTest.scala deleted file mode 100644 index 440faafc6f..0000000000 --- a/util-core/src/test/scala/com/twitter/io/exp/ActivitySourceTest.scala +++ /dev/null @@ -1,104 +0,0 @@ -package com.twitter.io.exp - -import com.twitter.io.Buf -import com.twitter.util.{Activity, FuturePool, MockTimer, Time} -import java.io.{File, ByteArrayInputStream} -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.{BeforeAndAfter, FunSuite} -import scala.util.Random - -@RunWith(classOf[JUnitRunner]) -class ActivitySourceTest extends FunSuite with BeforeAndAfter { - - val ok = new ActivitySource[String] { - def get(varName: String) = Activity.value(varName) - } - - val failed = new ActivitySource[Nothing] { - def get(varName: String) = Activity.exception(new Exception(varName)) - } - - val tempFile = Random.alphanumeric.take(5).mkString - - def writeToTempFile(s: String): Unit = { - val printer = new java.io.PrintWriter(new File(tempFile)) - try { - printer.print(s) - } finally { - printer.close() - } - } - - def bufToString(buf: Buf): String = buf match { - case Buf.Utf8(s) => s - case _ => "" - } - - before { - writeToTempFile("foo bar") - } - - after { - new File(tempFile).delete() - } - - test("ActivitySource.orElse") { - val a = failed.orElse(ok) - val b = ok.orElse(failed) - - assert("42" == a.get("42").sample()) - assert("42" == b.get("42").sample()) - } - - test("CachingActivitySource") { - val cache = new CachingActivitySource[String](new ActivitySource[String] { - def get(varName: String) = Activity.value(Random.alphanumeric.take(10).mkString) - }) - - val a = cache.get("a") - assert(a == cache.get("a")) - } - - test("FilePollingActivitySource") { - Time.withCurrentTimeFrozen { timeControl => - import com.twitter.conversions.time._ - implicit val timer = new MockTimer - val file = new FilePollingActivitySource(1.microsecond, FuturePool.immediatePool) - val buf = file.get(tempFile) - - var content: Option[String] = None - - val listen = buf.run.changes.respond { - case Activity.Ok(b) => - content = Some(bufToString(b)) - case _ => - content = None - } - - timeControl.advance(2.microsecond) - timer.tick() - assert(content == Some("foo bar")) - - writeToTempFile("foo baz") - - timeControl.advance(2.microsecond) - timer.tick() - assert(content == Some("foo baz")) - - listen.close() - } - } - - test("ClassLoaderActivitySource") { - val classLoader = new ClassLoader() { - override def getResourceAsStream(name: String) = - new ByteArrayInputStream(name.getBytes("UTF-8")) - } - - val loader = new ClassLoaderActivitySource(classLoader, FuturePool.immediatePool) - val bufAct = loader.get("bar baz") - - assert("bar baz" == bufToString(bufAct.sample())) - } -} diff --git a/util-core/src/test/scala/com/twitter/io/exp/MinimumThroughputTest.scala b/util-core/src/test/scala/com/twitter/io/exp/MinimumThroughputTest.scala deleted file mode 100644 index 1eb09ff777..0000000000 --- a/util-core/src/test/scala/com/twitter/io/exp/MinimumThroughputTest.scala +++ /dev/null @@ -1,262 +0,0 @@ -package com.twitter.io.exp - -import com.twitter.conversions.time._ -import com.twitter.io.{Writer, Buf, BufReader, Reader} -import com.twitter.util._ -import java.io.ByteArrayOutputStream -import org.junit.runner.RunWith -import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar - -@RunWith(classOf[JUnitRunner]) -class MinimumThroughputTest extends FunSuite with MockitoSugar { - test("Reader - negative minBps") { - intercept[AssertionError] { - MinimumThroughput.reader(mock[Reader[Buf]], -1, Timer.Nil) - } - } - - test("Reader - faster than min") { - val buf = Buf.UsAscii("soylent green is made of...") // 27 bytes - val reader = MinimumThroughput.reader(BufReader(buf), 0d, Timer.Nil) - - // read from the beginning - Await.result(reader.read(13)) match { - case Some(b) => assert(b == Buf.UsAscii("soylent green")) - case _ => fail() - } - - // a no-op read - Await.result(reader.read(0)) match { - case Some(b) => assert(b == Buf.Empty) - case _ => fail() - } - - // read to the end - Await.result(reader.read(20)) match { - case Some(b) => assert(b == Buf.UsAscii(" is made of...")) - case _ => fail() - } - - // read past the end - Await.result(reader.read(20)) match { - case None => - case _ => fail() - } - } - - test("Reader - below threshold after read") { - val buf = Buf.UsAscii("0") - - Time.withCurrentTimeFrozen { tc => - // it'd be great if we could just use a mock. - // the problem was that stubs are evaluated once, eagerly, at creation time. - val underlying = new Reader[Buf] { - private var reads = 0 - def discard(): Unit = () - def read(n: Int): Future[Option[Buf]] = { - reads += 1 - if (reads == 2) { - // now take 10 seconds to read 1 more byte, - // which would put us at 0.2 bps, and thus below the threshold. - tc.advance(10.seconds) - } - Future.value(Some(buf)) - } - } - - val reader = MinimumThroughput.reader( - underlying, - 1d, // min bytes per second - Timer.Nil - ) - - // do a read of 1 byte in 0 time — which is ok. - Await.result(reader.read(1)) match { - case Some(b) => assert(b == buf.slice(0, 1)) - case _ => fail() - } - - val ex = intercept[BelowThroughputException] { - // note in the mock above, the 2nd read takes 10 seconds - Await.result(reader.read(1)) - } - assert(ex.elapsed == 10.seconds) - assert(ex.expectedBps == 1d) - assert(ex.currentBps == 0.2d) - } - } - - test("Reader - times out while reading") { - val underlying = mock[Reader[Buf]] - when(underlying.read(1)) - .thenReturn(Future.never) - - val timer = new MockTimer() - val reader = MinimumThroughput.reader( - underlying, - 1d, // min bytes per second - timer - ) - - Time.withCurrentTimeFrozen { tc => - val f = reader.read(1) - tc.advance(10.seconds) - timer.tick() - - val ex = intercept[BelowThroughputException] { - Await.result(f) - } - assert(ex.elapsed == 10.seconds) - assert(ex.expectedBps == 1d) - assert(ex.currentBps == 0d) - } - } - - test("Reader - failures from underlying reader are untouched") { - val ex = new RuntimeException("└[∵┌]└[ ∵ ]┘[┐∵]┘") - val underlying = mock[Reader[Buf]] - when(underlying.read(1)) - .thenReturn(Future.exception(ex)) - - val reader = MinimumThroughput.reader( - underlying, - 1d, // min bytes per second - Timer.Nil - ) - - val thrown = intercept[RuntimeException] { - Await.result(reader.read(1)) - } - assert(thrown == ex) - } - - test("Reader - pass through EOFs from underlying") { - val reader = MinimumThroughput.reader( - Reader.empty, - 1d, // min bytes per second - Timer.Nil - ) - - Await.result(reader.read(1)) match { - case None => - case _ => fail() - } - } - - test("Reader - discard is passed through to underlying") { - val underlying = mock[Reader[Buf]] - val reader = MinimumThroughput.reader(underlying, 1, Timer.Nil) - - reader.discard() - verify(underlying).discard() - } - - test("Writer - faster than min") { - val buf = Buf.UsAscii("0") - - val writer = - MinimumThroughput.writer(Writer.fromOutputStream(new ByteArrayOutputStream()), 0d, Timer.Nil) - - val w1 = writer.write(buf) - Await.ready(w1) - assert(w1.isDone) - - val w2 = writer.write(buf) - Await.ready(w2) - assert(w2.isDone) - } - - test("Writer - below threshold after write") { - val buf = Buf.UsAscii("0") - - Time.withCurrentTimeFrozen { tc => - // it'd be great if we could just use a mock. - // the problem was that stubs are evaluated once, eagerly, at creation time. - val underlying = new Writer[Buf] { - private var writes = 0 - def fail(cause: Throwable): Unit = () - def write(buf: Buf): Future[Unit] = { - writes += 1 - if (writes == 2) { - // now take 10 seconds to write 1 more byte, - // which would put us at 0.2 bps, and thus below the threshold. - tc.advance(10.seconds) - } - Future.Done - } - } - - val writer = MinimumThroughput.writer( - underlying, - 1d, // min bytes per second - Timer.Nil - ) - - // do a write of 1 byte in 0 time — which is ok. - val w1 = writer.write(buf) - Await.ready(w1) - assert(w1.isDone) - - val ex = intercept[BelowThroughputException] { - // note in the mock above, the 2nd write takes 10 seconds - Await.result(writer.write(buf)) - } - assert(ex.elapsed == 10.seconds) - assert(ex.expectedBps == 1d) - assert(ex.currentBps == 0.2d) - } - } - - test("Writer - times out while writing") { - val buf = Buf.UsAscii("0") - val underlying = mock[Writer[Buf]] - when(underlying.write(buf)) - .thenReturn(Future.never) - - val timer = new MockTimer() - val writer = MinimumThroughput.writer(underlying, 1d, timer) - - Time.withCurrentTimeFrozen { tc => - val f = writer.write(buf) - tc.advance(10.seconds) - timer.tick() - - val ex = intercept[BelowThroughputException] { - Await.result(f) - } - assert(ex.elapsed == 10.seconds) - assert(ex.expectedBps == 1d) - assert(ex.currentBps == 0d) - } - } - - test("Writer - failures from underlying writer are untouched") { - val buf = Buf.UsAscii("0") - val ex = new RuntimeException("ᕕ( ᐛ )ᕗ") - val underlying = mock[Writer[Buf]] - when(underlying.write(buf)) - .thenReturn(Future.exception(ex)) - - val writer = MinimumThroughput.writer(underlying, 1d, Timer.Nil) - - val thrown = intercept[RuntimeException] { - Await.result(writer.write(buf)) - } - assert(thrown == ex) - } - - test("Writer - fail is passed through to underlying") { - val underlying = mock[Writer[Buf]] - - val writer = MinimumThroughput.writer(underlying, 1d, Timer.Nil) - - val ex = new RuntimeException() - writer.fail(ex) - - verify(underlying).fail(ex) - } - -} diff --git a/util-core/src/test/scala/com/twitter/util/ActivityTest.scala b/util-core/src/test/scala/com/twitter/util/ActivityTest.scala index 2509d10519..4bec850cb5 100644 --- a/util-core/src/test/scala/com/twitter/util/ActivityTest.scala +++ b/util-core/src/test/scala/com/twitter/util/ActivityTest.scala @@ -1,12 +1,10 @@ package com.twitter.util +import com.twitter.conversions.DurationOps._ import java.util.concurrent.atomic.AtomicReference -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ActivityTest extends FunSuite { +class ActivityTest extends AnyFunSuite { test("Activity#flatMap") { val v = Var(Activity.Pending: Activity.State[Int]) val ref = new AtomicReference[Seq[Activity.State[Int]]] @@ -87,24 +85,24 @@ class ActivityTest extends FunSuite { } test("Activity.future: produce an initially-pending Activity") { - assert(Activity.future(Future.never).run.sample == Activity.Pending) + assert(Activity.future(Future.never).run.sample() == Activity.Pending) } test("Activity.future: produce an Activity that completes on success of the original Future") { val p = new Promise[Int] val act = Activity.future(p) - assert(act.run.sample == Activity.Pending) + assert(act.run.sample() == Activity.Pending) p.setValue(4) - assert(act.run.sample == Activity.Ok(4)) + assert(act.run.sample() == Activity.Ok(4)) } test("Activity.future: produce an Activity that fails on failure of the original Future") { val p = new Promise[Unit] val e = new Exception("gooby pls") val act = Activity.future(p) - assert(act.run.sample == Activity.Pending) + assert(act.run.sample() == Activity.Pending) p.setException(e) - assert(act.run.sample == Activity.Failed(e)) + assert(act.run.sample() == Activity.Failed(e)) } test( @@ -113,7 +111,7 @@ class ActivityTest extends FunSuite { ) { val p = new Promise[Unit] val obs = Activity.future(p).run.changes.register(Witness((_: Any) => ())) - Await.ready(obs.close()) + Await.result(obs.close(), 5.seconds) assert(!p.isDefined) } @@ -200,47 +198,182 @@ class ActivityTest extends FunSuite { ) } - test("Activity.stabilize") { + test("Activity#stabilize will update until an Ok event") { val ex = new Exception val v = Var[Activity.State[Int]](Activity.Pending) + val a = Activity(v) - val w = Witness(v) val ref = new AtomicReference[Seq[Activity.State[Int]]] a.stabilize.states.build.register(Witness(ref)) assert(ref.get == Seq(Activity.Pending)) - w.notify(Activity.Failed(ex)) + v() = Activity.Failed(ex) assert(ref.get == Seq(Activity.Pending, Activity.Failed(ex))) - w.notify(Activity.Ok(1)) + v() = Activity.Ok(1) assert(ref.get == Seq(Activity.Pending, Activity.Failed(ex), Activity.Ok(1))) + } - w.notify(Activity.Failed(ex)) - assert(ref.get == Seq(Activity.Pending, Activity.Failed(ex), Activity.Ok(1), Activity.Ok(1))) + test("Activity#stabilize will stop updating on non-Oks after an Ok event") { + val ex = new Exception + val v = Var[Activity.State[Int]](Activity.Pending) - w.notify(Activity.Pending) - assert( - ref.get == Seq( - Activity.Pending, - Activity.Failed(ex), - Activity.Ok(1), - Activity.Ok(1), - Activity.Ok(1) - ) - ) + val a = Activity(v) - w.notify(Activity.Ok(2)) - assert( - ref.get == Seq( - Activity.Pending, - Activity.Failed(ex), - Activity.Ok(1), - Activity.Ok(1), - Activity.Ok(1), - Activity.Ok(2) - ) - ) + val ref = new AtomicReference[Seq[Activity.State[Int]]] + a.stabilize.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Pending)) + + v() = Activity.Ok(1) + assert(ref.get == Seq(Activity.Pending, Activity.Ok(1))) + + v() = Activity.Failed(ex) + assert(ref.get == Seq(Activity.Pending, Activity.Ok(1), Activity.Ok(1))) + + v() = Activity.Pending + assert(ref.get == Seq(Activity.Pending, Activity.Ok(1), Activity.Ok(1), Activity.Ok(1))) + } + + test("Activity#stabilize will also keep the old event on non-Ok after it starts in Ok") { + val ex = new Exception + val v = Var[Activity.State[Int]](Activity.Ok(1)) + + val a = Activity(v) + + val ref = new AtomicReference[Seq[Activity.State[Int]]] + a.stabilize.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Ok(1))) + + v() = Activity.Failed(ex) + assert(ref.get == Seq(Activity.Ok(1), Activity.Ok(1))) + } + + test("Activity#stabilize can still transition to a different Ok") { + val v = Var[Activity.State[Int]](Activity.Ok(1)) + val a = Activity(v) + + val ref = new AtomicReference[Seq[Activity.State[Int]]] + a.stabilize.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Ok(1))) + + v() = Activity.Ok(2) + assert(ref.get == Seq(Activity.Ok(1), Activity.Ok(2))) + } + + test("Activity#stabilize will keep the old stable value when not observed") { + val ex = new Exception + val v = Var[Activity.State[Int]](Activity.Ok(1)) + val a = Activity(v).stabilize + + val ref = new AtomicReference[Seq[Activity.State[Int]]] + val c = a.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Ok(1))) + Await.result(c.close(), 5.seconds) + + assert(a.sample() == 1) + } + + test( + "Activity#stabilize will ignore new values until it is observed again after deregistration") { + val ex = new Exception + val v = Var[Activity.State[Int]](Activity.Ok(1)) + val a = Activity(v).stabilize + + val ref = new AtomicReference[Seq[Activity.State[Int]]] + val c = a.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Ok(1))) + Await.result(c.close(), 5.seconds) + + assert(a.sample() == 1) + + // this value does not get cached + v() = Activity.Ok(2) + + // this is the true value of the underlying item + v() = Activity.Pending + + // uses the cached value + assert(a.sample() == 1) + } + + test( + "Activity#stabilize will see new valid values immediately on observation after deregistration") { + val ex = new Exception + val v = Var[Activity.State[Int]](Activity.Ok(1)) + val a = Activity(v).stabilize + + val ref = new AtomicReference[Seq[Activity.State[Int]]] + val c = a.states.build.register(Witness(ref)) + + assert(ref.get == Seq(Activity.Ok(1))) + Await.result(c.close(), 5.seconds) + + assert(a.sample() == 1) + + v() = Activity.Ok(2) + assert(a.sample() == 2) + } + + class FakeEvent extends Event[Activity.State[Int]] { + var count = 0 + var w: Witness[Activity.State[Int]] = null + + def register(witness: Witness[Activity.State[Int]]): Closable = { + count += 1 + w = witness + Closable.make { (deadline: Time) => + count -= 1 + w = null + Future.Done + } + } + } + + test("Activity.apply(Event) doesn't leak when sampled") { + val evt = new FakeEvent + val activity = Activity(evt) + assert(evt.count == 0) + + assert(activity.run.sample() == Activity.Pending) + assert(evt.count == 0) + } + + test("Activity.apply(Event) doesn't leak with derived events") { + val evt = new FakeEvent + val activity = Activity(evt) + assert(evt.count == 0) + + val derivedEvt = activity.run.changes + assert(evt.count == 0) + val ref = new AtomicReference[Activity.State[Int]]() + val closable = derivedEvt.register(Witness(ref)) + assert(evt.count == 1) + assert(ref.get == Activity.Pending) + evt.w.notify(Activity.Ok(1)) + assert(Activity.sample(activity) == 1) + assert(evt.count == 1) + Await.result(closable.close(), 5.seconds) + assert(evt.count == 0) + } + + test("Activity.apply(Event) doesn't preserve events after it stops watching") { + val evt = new FakeEvent + val activity = Activity(evt) + + val derivedEvt = activity.run.changes + val ref = new AtomicReference[Activity.State[Int]]() + val closable = derivedEvt.register(Witness(ref)) + evt.w.notify(Activity.Ok(1)) + assert(Activity.sample(activity) == 1) + Await.result(closable.close(), 5.seconds) + + assert(activity.run.sample() == Activity.Pending) } } diff --git a/util-core/src/test/scala/com/twitter/util/BUILD b/util-core/src/test/scala/com/twitter/util/BUILD new file mode 100644 index 0000000000..7824c17828 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/util/BUILD @@ -0,0 +1,22 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/test/java/com/twitter/util:object_size_calculator", + ], +) diff --git a/util-core/src/test/scala/com/twitter/util/Base64LongTest.scala b/util-core/src/test/scala/com/twitter/util/Base64LongTest.scala index 09fb65a9ee..a8c3422153 100644 --- a/util-core/src/test/scala/com/twitter/util/Base64LongTest.scala +++ b/util-core/src/test/scala/com/twitter/util/Base64LongTest.scala @@ -1,13 +1,11 @@ package com.twitter.util import com.twitter.util.Base64Long.{fromBase64, StandardBase64Alphabet, toBase64} -import org.junit.runner.RunWith import org.scalacheck.Arbitrary import org.scalacheck.Arbitrary.arbitrary import org.scalacheck.Gen -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite object Base64LongTest { private val BoundaryValues = @@ -29,7 +27,7 @@ object Base64LongTest { go(Nil, 64, 0) } - private implicit val arbAlphabet = + private implicit val arbAlphabet: Arbitrary[Alphabet] = Arbitrary[Alphabet]( Gen.oneOf( Gen.const(StandardBase64Alphabet), @@ -54,8 +52,7 @@ object Base64LongTest { } } -@RunWith(classOf[JUnitRunner]) -class Base64LongTest extends FunSuite with GeneratorDrivenPropertyChecks { +class Base64LongTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { import Base64LongTest._ private[this] def forAllLongs(f: (Alphabet, Long) => Unit): Unit = { @@ -80,9 +77,7 @@ class Base64LongTest extends FunSuite with GeneratorDrivenPropertyChecks { } test("toBase64 uses the expected number of digits") { - BoundaryValues.foreach { (n: Long) => - assert(toBase64(n).length == expectedLength(n)) - } + BoundaryValues.foreach { (n: Long) => assert(toBase64(n).length == expectedLength(n)) } forAll((n: Long) => assert(toBase64(n).length == expectedLength(n))) val b = new StringBuilder @@ -94,9 +89,7 @@ class Base64LongTest extends FunSuite with GeneratorDrivenPropertyChecks { } test("fromBase64 is the inverse of toBase64") { - BoundaryValues.foreach { l => - assert(l == fromBase64(toBase64(l))) - } + BoundaryValues.foreach { l => assert(l == fromBase64(toBase64(l))) } val b = new StringBuilder forAllLongs { (a, n) => @@ -115,8 +108,11 @@ class Base64LongTest extends FunSuite with GeneratorDrivenPropertyChecks { test("fromBase64 throws an IllegalArgumentException exception for characters out of range") { forAll { (s: String, a: Alphabet) => - if (s.exists(!a.isDefinedAt(_))) { - assertThrows[IllegalArgumentException](fromBase64(s)) + val inverted = Base64Long.invertAlphabet(a) + if (s.exists(!inverted.isDefinedAt(_))) { + assertThrows[IllegalArgumentException]( + fromBase64(s, 0, s.length, inverted) + ) } } } diff --git a/util-core/src/test/scala/com/twitter/util/BoundedStackTest.scala b/util-core/src/test/scala/com/twitter/util/BoundedStackTest.scala deleted file mode 100644 index 1954326f87..0000000000 --- a/util-core/src/test/scala/com/twitter/util/BoundedStackTest.scala +++ /dev/null @@ -1,105 +0,0 @@ -package com.twitter.util - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class BoundedStackTest extends WordSpec { - "BoundedStack" should { - "empty" in { - val buf = new BoundedStack[String](4) - assert(buf.length == 0) - assert(buf.size == 0) - assert(buf.isEmpty == true) - intercept[IndexOutOfBoundsException] { - buf(0) - } - intercept[NoSuchElementException] { - buf.pop - } - assert(buf.iterator.hasNext == false) - } - - "handle single element" in { - val buf = new BoundedStack[String](4) - buf += "a" - assert(buf.size == 1) - assert(buf(0) == "a") - assert(buf.toList == List("a")) - } - - "handle multiple element" in { - val buf = new BoundedStack[String](4) - buf ++= List("a", "b", "c") - assert(buf.size == 3) - assert(buf(0) == "c") - assert(buf(1) == "b") - assert(buf(2) == "a") - assert(buf.toList == List("c", "b", "a")) - assert(buf.pop == "c") - assert(buf.size == 2) - assert(buf.pop == "b") - assert(buf.size == 1) - assert(buf.pop == "a") - assert(buf.size == 0) - } - - "handle overwrite/rollover" in { - val buf = new BoundedStack[String](4) - buf ++= List("a", "b", "c", "d", "e", "f") - assert(buf.size == 4) - assert(buf(0) == "f") - assert(buf.toList == List("f", "e", "d", "c")) - } - - "handle update" in { - val buf = new BoundedStack[String](4) - buf ++= List("a", "b", "c", "d", "e", "f") - for (i <- 0 until buf.size) { - val old = buf(i) - val updated = old + "2" - buf(i) = updated - assert(buf(i) == updated) - } - assert(buf.toList == List("f2", "e2", "d2", "c2")) - } - - "insert at 0 is same as +=" in { - val buf = new BoundedStack[String](3) - buf.insert(0, "a") - assert(buf.size == 1) - assert(buf(0) == "a") - buf.insert(0, "b") - assert(buf.size == 2) - assert(buf(0) == "b") - assert(buf(1) == "a") - buf.insert(0, "c") - assert(buf(0) == "c") - assert(buf(1) == "b") - assert(buf(2) == "a") - buf.insert(0, "d") - } - - "insert at count pushes to bottom" in { - val buf = new BoundedStack[String](3) - buf.insert(0, "a") - buf.insert(1, "b") - buf.insert(2, "c") - assert(buf(0) == "a") - assert(buf(1) == "b") - assert(buf(2) == "c") - } - - "insert > count throws exception" in { - val buf = new BoundedStack[String](3) - intercept[IndexOutOfBoundsException] { - buf.insert(1, "a") - } - buf.insert(0, "a") - intercept[IndexOutOfBoundsException] { - buf.insert(2, "b") - } - } - } -} diff --git a/util-core/src/test/scala/com/twitter/util/CancellableTest.scala b/util-core/src/test/scala/com/twitter/util/CancellableTest.scala index 0d84b5dd79..8b23f172f9 100644 --- a/util-core/src/test/scala/com/twitter/util/CancellableTest.scala +++ b/util-core/src/test/scala/com/twitter/util/CancellableTest.scala @@ -1,19 +1,70 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class CancellableTest extends WordSpec { - "CancellableSink" should { - "cancel once" in { - var count = 0 - val s = new CancellableSink { count += 1 } - s.cancel() - assert(count == 1) - s.cancel() - assert(count == 1) +import org.scalatest.funsuite.AnyFunSuite + +class CancellableTest extends AnyFunSuite { + + test("Cancellable.nil is always cancelled") { + assert(!Cancellable.nil.isCancelled) + } + + test("Cancellable.nil cancel method") { + Cancellable.nil.cancel() + assert(!Cancellable.nil.isCancelled) + } + + test("Cancellable.nil linkTo method") { + Cancellable.nil.linkTo(Cancellable.nil) + assert(!Cancellable.nil.isCancelled) + } + + test("CancellableSink cancel normal case") { + var count = 2 + def square(): Unit = count *= count + val s = new CancellableSink(square()) + assert(count == 2) + s.cancel() + assert(count == 4) + } + + test("CancellableSink cancel once") { + var count = 0 + def increment(): Unit = count += 1 + val s = new CancellableSink(increment()) + s.cancel() + assert(count == 1) + s.cancel() + assert(count == 1) + } + + test("CancellableSink has not cancelled") { + var count = 0 + def increment(): Unit = count += 1 + val s = new CancellableSink(increment()) + assert(!s.isCancelled) + assert(count == 0) + } + + test("CancellableSink confirm cancelled") { + var count = 0 + def increment(): Unit = count += 1 + val s = new CancellableSink(increment()) + s.cancel() + assert(s.isCancelled) + assert(count == 1) + } + + test("CancellableSink not support linking") { + var count = 1 + def multiply(): Unit = count *= 2 + val s1 = new CancellableSink(multiply()) + val s2 = new CancellableSink(multiply()) + assertThrows[Exception] { + s1.linkTo(s2) + } + assertThrows[Exception] { + s2.linkTo(s1) } } + } diff --git a/util-core/src/test/scala/com/twitter/util/ClosableOnceTest.scala b/util-core/src/test/scala/com/twitter/util/ClosableOnceTest.scala index 8801fb91f0..2d234bf8bf 100644 --- a/util-core/src/test/scala/com/twitter/util/ClosableOnceTest.scala +++ b/util-core/src/test/scala/com/twitter/util/ClosableOnceTest.scala @@ -1,9 +1,9 @@ package com.twitter.util -import com.twitter.conversions.time._ -import org.scalatest.FunSuite +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite -class ClosableOnceTest extends FunSuite { +class ClosableOnceTest extends AnyFunSuite { private def ready[T <: Awaitable[_]](awaitable: T): T = Await.ready(awaitable, 10.seconds) @@ -18,7 +18,9 @@ class ClosableOnceTest extends FunSuite { } val closableOnce = ClosableOnce.of(underlying) + assert(closableOnce.isClosed == false) closableOnce.close() + assert(closableOnce.isClosed == true) closableOnce.close() assert(closedCalls == 1) } @@ -33,8 +35,11 @@ class ClosableOnceTest extends FunSuite { } } + assert(closableOnce.isClosed == false) + assert(ready(closableOnce.close()).poll.get.throwable == ex) assert(closedCalls == 1) + assert(closableOnce.isClosed == true) assert(ready(closableOnce.close()).poll.get.throwable == ex) assert(closedCalls == 1) diff --git a/util-core/src/test/scala/com/twitter/util/ClosableTest.scala b/util-core/src/test/scala/com/twitter/util/ClosableTest.scala index e74b9ef3b7..20c4cd3d41 100644 --- a/util-core/src/test/scala/com/twitter/util/ClosableTest.scala +++ b/util-core/src/test/scala/com/twitter/util/ClosableTest.scala @@ -1,14 +1,34 @@ package com.twitter.util -import com.twitter.util.TimeConversions.intToTimeableNumber -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.concurrent.{Eventually, IntegrationPatience} -import org.scalatest.junit.JUnitRunner -import scala.language.reflectiveCalls - -@RunWith(classOf[JUnitRunner]) -class ClosableTest extends FunSuite with Eventually with IntegrationPatience { +import com.twitter.conversions.DurationOps._ +import org.scalatest.concurrent.Eventually +import org.scalatest.concurrent.IntegrationPatience +import org.scalatest.funsuite.AnyFunSuite + +class ClosableTest extends AnyFunSuite with Eventually with IntegrationPatience { + + // Workaround methods for dealing with Scala compiler warnings: + // + // assert(f.isDone) cannot be called directly because it results in a Scala compiler + // warning: 'possible missing interpolator: detected interpolated identifier `$conforms`' + // + // This is due to the implicit evidence required for `Future.isDone` which checks to see + // whether the Future that is attempting to have `isDone` called on it conforms to the type + // of `Future[Unit]`. This is done using `Predef.$conforms` + // https://www.scala-lang.org/api/2.12.2/scala/Predef$.html#$conforms[A]:A%3C:%3CA + // + // Passing that directly to `assert` causes problems because the `$conforms` is also seen as + // an interpolated string. We get around it by evaluating first and passing the result to + // `assert`. + private[this] def isDone(f: Future[Unit]): Boolean = + f.isDefined + + private[this] def assertIsDone(f: Future[Unit]): Unit = + assert(isDone(f)) + + private[this] def assertIsNotDone(f: Future[Unit]): Unit = + assert(!isDone(f)) + test("Closable.close(Duration)") { Time.withCurrentTimeFrozen { _ => var time: Option[Time] = None @@ -49,16 +69,17 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { assert(n1 == 1) assert(n2 == 1) - assert(!f.isDone) + assertIsNotDone(f) p1.setDone() - assert(!f.isDone) + assertIsNotDone(f) p2.setDone() - assert(f.isDone) + assertIsDone(f) } test("Closable.all with exceptions") { val throwing = Closable.make(_ => sys.error("lolz")) - val tracking = new Closable { + + class TrackedClosable extends Closable { @volatile var calledClose = false override def close(deadline: Time): Future[Unit] = { @@ -67,6 +88,7 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { } } + val tracking = new TrackedClosable() val f = Closable.all(throwing, tracking).close() intercept[Exception] { Await.result(f, 5.seconds) @@ -78,36 +100,63 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { assert( Future .value(1) - .map(_ => Closable.all(Closable.nop, Closable.nop).close().isDone) + .map(_ => Closable.all(Closable.nop, Closable.nop).close().isDefined) .poll .contains(Return.True) ) } + test("Cloasable.make catches NonFatals and translates to failed Futures") { + val throwing = Closable.make { _ => throw new Exception("lolz") } + + var failed = 0 + val f = throwing.close().respond { + case Throw(_) => failed = -1 + case _ => failed = 1 + } + + val ex = intercept[Exception] { + Await.result(f, 2.seconds) + } + assert(ex.getMessage.contains("lolz")) + assert(failed == -1) + } +} + +abstract class ClosableSequenceTest extends AnyFunSuite with Eventually with IntegrationPatience { + + def sequenceTwo(a: Closable, b: Closable): Closable + + private[this] def assertIsDone(f: Future[Unit]): Unit = + assert(f.isDefined) + + private[this] def assertIsNotDone(f: Future[Unit]): Unit = + assert(!f.isDefined) + test("Closable.sequence") { val p1, p2 = new Promise[Unit] var n1, n2 = 0 val c1 = Closable.make(_ => { n1 += 1; p1 }) val c2 = Closable.make(_ => { n2 += 1; p2 }) - val c = Closable.sequence(c1, c2) + val c = sequenceTwo(c1, c2) assert(n1 == 0) assert(n2 == 0) val f = c.close() assert(n1 == 1) assert(n2 == 0) - assert(!f.isDone) + assertIsNotDone(f) p1.setDone() assert(n1 == 1) assert(n2 == 1) - assert(!f.isDone) + assertIsNotDone(f) p2.setDone() assert(n1 == 1) assert(n2 == 1) - assert(f.isDone) + assertIsDone(f) } test("Closable.sequence is resilient to failed closes") { @@ -122,7 +171,7 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { } // test with a failed close first - val f1 = Closable.sequence(c1, c2).close() + val f1 = sequenceTwo(c1, c2).close() assert(n1 == 1) assert(n2 == 1) val x1 = intercept[RuntimeException] { @@ -131,7 +180,7 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { assert(x1.getMessage.contains("n1=1")) // then test in reverse order, failed close last - val f2 = Closable.sequence(c2, c1).close() + val f2 = sequenceTwo(c2, c1).close() assert(n1 == 2) assert(n2 == 2) val x2 = intercept[RuntimeException] { @@ -139,19 +188,20 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { } assert(x2.getMessage.contains("n1=2")) - // multiple failures, returns the first failure - val f3 = Closable.sequence(c1, c1).close() + // multiple failures, returns the first failure, and a list of suppressed failures. + val f3 = sequenceTwo(c1, c1).close() assert(n1 == 4) val x3 = intercept[RuntimeException] { Await.result(f3, 5.seconds) } assert(x3.getMessage.contains("n1=3")) + assert(x3.getSuppressed.head.getMessage.contains("n1=4")) // verify that first failure is returned, but only after the // last closable is satisfied val p = new Promise[Unit]() val c3 = Closable.make(_ => p) - val f4 = Closable.sequence(c1, c3).close() + val f4 = sequenceTwo(c1, c3).close() assert(n1 == 5) assert(!f4.isDefined) p.setDone() @@ -163,7 +213,7 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { test("Closable.sequence lifts synchronous exceptions") { val throwing = Closable.make(_ => sys.error("lolz")) - val f = Closable.sequence(Closable.nop, throwing).close() + val f = sequenceTwo(Closable.nop, throwing).close() val ex = intercept[Exception] { Await.result(f, 5.seconds) } @@ -174,27 +224,17 @@ class ClosableTest extends FunSuite with Eventually with IntegrationPatience { assert( Future .value(1) - .map(_ => Closable.sequence(Closable.nop, Closable.nop).close().isDone) + .map(_ => sequenceTwo(Closable.nop, Closable.nop).close().isDefined) .poll .contains(Return.True) ) } +} - test("Cloasable.make catches NonFatals and translates to failed Futures") { - val throwing = Closable.make { _ => - throw new Exception("lolz") - } - - var failed = 0 - val f = throwing.close().respond { - case Throw(_) => failed = -1 - case _ => failed = 1 - } +final class ClosableSequence2Test extends ClosableSequenceTest { + def sequenceTwo(a: Closable, b: Closable): Closable = Closable.sequence(a, b) +} - val ex = intercept[Exception] { - Await.result(f, 2.seconds) - } - assert(ex.getMessage.contains("lolz")) - assert(failed == -1) - } +final class ClosableSequenceNTest extends ClosableSequenceTest { + def sequenceTwo(a: Closable, b: Closable): Closable = Closable.sequence(Seq(a, b): _*) } diff --git a/util-core/src/test/scala/com/twitter/util/CloseAwaitablyTest.scala b/util-core/src/test/scala/com/twitter/util/CloseAwaitablyTest.scala index b5ea3ed1da..0aea96b7e6 100644 --- a/util-core/src/test/scala/com/twitter/util/CloseAwaitablyTest.scala +++ b/util-core/src/test/scala/com/twitter/util/CloseAwaitablyTest.scala @@ -1,15 +1,12 @@ package com.twitter.util -import com.twitter.conversions.time._ -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class CloseAwaitablyTest extends FunSuite { +class CloseAwaitablyTest extends AnyFunSuite { class Context extends Closable with CloseAwaitably { - val p = new Promise[Unit] - var n = 0 + val p: Promise[Unit] = new Promise[Unit] + var n: Int = 0 def close(deadline: Time): Future[Unit] = closeAwaitably { n += 1 p @@ -55,9 +52,7 @@ class CloseAwaitablyTest extends FunSuite { test("close awaitably with error") { val message = "FORCED EXCEPTION" - val closeFn: () => Future[Unit] = { () => - throw new Exception(message) - } + val closeFn: () => Future[Unit] = { () => throw new Exception(message) } val testClosable = new TestClosable(closeFn) val f = testClosable.close(Time.now) @@ -71,9 +66,7 @@ class CloseAwaitablyTest extends FunSuite { test("close awaitably with fatal error, fatal error blows up close()") { val message = "FORCED INTERRUPTED EXCEPTION" - val closeFn: () => Future[Unit] = { () => - throw new InterruptedException(message) - } + val closeFn: () => Future[Unit] = { () => throw new InterruptedException(message) } val testClosable = new TestClosable(closeFn) val e = intercept[InterruptedException] { diff --git a/util-core/src/test/scala/com/twitter/util/ConfigTest.scala b/util-core/src/test/scala/com/twitter/util/ConfigTest.scala deleted file mode 100644 index 70749bde78..0000000000 --- a/util-core/src/test/scala/com/twitter/util/ConfigTest.scala +++ /dev/null @@ -1,91 +0,0 @@ -package com.twitter.util - -import org.junit.runner.RunWith -import org.scalatest.{Matchers, WordSpec} -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class ConfigTest extends WordSpec with Matchers { - import Config._ - - "Config" should { - "computed should delay evaluation" in { - class Foo extends Config.Nothing { - var didIt = false - var x = 10 - var y = computed { - didIt = true - x * 2 + 5 - } - } - - val foo = new Foo - assert(!foo.didIt) - assert((foo.y: Int) == 25) // use type annotation to force implicit conversion - assert(foo.didIt) - } - - "subclass can override indepedent var for use in dependent var" in { - class Foo extends Config.Nothing { - var x = 10 - var y = computed(x * 2 + 5) - } - val bar = new Foo { - x = 20 - } - assert((bar.y: Int) == 45) // use type annotation to force implicit conversion - } - - class Bar extends Config.Nothing { - var z = required[Int] - } - class Baz extends Config.Nothing { - var w = required[Int] - } - class Foo extends Config.Nothing { - var x = required[Int] - var y = 3 - var bar = required[Bar] - var baz = optional[Baz] - } - - "missingValues" should { - "must return empty Seq when no values are missing" in { - val foo = new Foo { - x = 42 - bar = new Bar { - z = 10 - } - } - assert(foo.missingValues.isEmpty) - } - - "must find top-level missing values" in { - val foo = new Foo - assert(foo.missingValues.sorted == Seq("x", "bar").sorted) - } - - "must find top-level and nested missing values" in { - val foo = new Foo { - bar = new Bar - } - assert(foo.missingValues.sorted == Seq("x", "bar.z").sorted) - } - - "must find nested missing values in optional sub-configs" in { - val foo = new Foo { - x = 3 - bar = new Bar { - z = 1 - } - baz = new Baz - } - // we can't do an exact match here since in scala 2.12.x - // this looks something like `List("$anonfun$new$16.w", "baz.w")`. - // this code is deprecated and we want to stop supporting it. - // so we simplify the test. - assert(foo.missingValues.contains("baz.w")) - } - } - } -} diff --git a/util-core/src/test/scala/com/twitter/util/DiffTest.scala b/util-core/src/test/scala/com/twitter/util/DiffTest.scala index ecc577275e..197c078436 100644 --- a/util-core/src/test/scala/com/twitter/util/DiffTest.scala +++ b/util-core/src/test/scala/com/twitter/util/DiffTest.scala @@ -1,13 +1,10 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner import org.scalacheck.Arbitrary.arbitrary -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class DiffTest extends FunSuite with GeneratorDrivenPropertyChecks { +class DiffTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { val f: Int => String = _.toString test("Diffable.ofSet") { diff --git a/util-core/src/test/scala/com/twitter/util/DurationTest.scala b/util-core/src/test/scala/com/twitter/util/DurationTest.scala index 22e551c230..c1ead91247 100644 --- a/util-core/src/test/scala/com/twitter/util/DurationTest.scala +++ b/util-core/src/test/scala/com/twitter/util/DurationTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,281 +16,284 @@ package com.twitter.util -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import java.util.concurrent.TimeUnit -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class DurationTest extends { val ops = Duration } with TimeLikeSpec[Duration] { - - "Duration" should { - "*" in { - assert(1.second * 2 == 2.seconds) - assert(1.second * Long.MaxValue == Duration.Top) - assert(1.second * Long.MinValue == Duration.Bottom) - assert(-10.seconds * Long.MinValue == Duration.Top) - assert(-10.seconds * Long.MaxValue == Duration.Bottom) - assert(500.milliseconds * 4 == 2.seconds) - assert(1.nanosecond * Long.MaxValue == Long.MaxValue.nanoseconds) - assert(1.nanosecond * Long.MinValue == Long.MinValue.nanoseconds) - - assert(1.second * 2.0 == 2.seconds) - assert(500.milliseconds * 4.0 == 2.seconds) - assert(1.second * 0.5 == 500.milliseconds) - assert(1.second * Double.PositiveInfinity == Duration.Top) - assert(1.second * Double.NegativeInfinity == Duration.Bottom) - assert(1.second * Double.NaN == Duration.Undefined) - assert(1.nanosecond * (Long.MaxValue.toDouble + 1) == Duration.Top) - assert(1.nanosecond * (Long.MinValue.toDouble - 1) == Duration.Bottom) - - assert(Duration.Undefined * Double.NaN == Duration.Undefined) - assert(Duration.Top * Double.NaN == Duration.Undefined) - assert(Duration.Bottom * Double.NaN == Duration.Undefined) - - forAll { (a: Long, b: Long) => - val c = BigInt(a) * BigInt(b) - val d = a.nanoseconds * b - - assert( - (c >= Long.MaxValue && d == Duration.Top) || - (c <= Long.MinValue && d == Duration.Bottom) || - (d == c.toLong.nanoseconds) - ) - } +import org.scalatest.funsuite.AnyFunSuite + +trait DurationOps { val ops: Duration.type = Duration } +class DurationTest extends AnyFunSuite with DurationOps with TimeLikeSpec[Duration] { + test("Duration *") { + assert(1.second * 2 == 2.seconds) + assert(1.second * Long.MaxValue == Duration.Top) + assert(1.second * Long.MinValue == Duration.Bottom) + assert(-10.seconds * Long.MinValue == Duration.Top) + assert(-10.seconds * Long.MaxValue == Duration.Bottom) + assert(500.milliseconds * 4 == 2.seconds) + assert(1.nanosecond * Long.MaxValue == Long.MaxValue.nanoseconds) + assert(1.nanosecond * Long.MinValue == Long.MinValue.nanoseconds) + + assert(1.second * 2.0 == 2.seconds) + assert(500.milliseconds * 4.0 == 2.seconds) + assert(1.second * 0.5 == 500.milliseconds) + assert(1.second * Double.PositiveInfinity == Duration.Top) + assert(1.second * Double.NegativeInfinity == Duration.Bottom) + assert(1.second * Double.NaN == Duration.Undefined) + assert(1.nanosecond * (Long.MaxValue.toDouble + 1) == Duration.Top) + assert(1.nanosecond * (Long.MinValue.toDouble - 1) == Duration.Bottom) + + assert(Duration.Undefined * Double.NaN == Duration.Undefined) + assert(Duration.Top * Double.NaN == Duration.Undefined) + assert(Duration.Bottom * Double.NaN == Duration.Undefined) + + forAll { (a: Long, b: Long) => + val c = BigInt(a) * BigInt(b) + val d = a.nanoseconds * b + + assert( + (c >= Long.MaxValue && d == Duration.Top) || + (c <= Long.MinValue && d == Duration.Bottom) || + (d == c.toLong.nanoseconds) + ) } + } - "/" in { - assert(10.seconds / 2 == 5.seconds) - assert(1.seconds / 0 == Duration.Top) - assert((-1).seconds / 0 == Duration.Bottom) - assert(0.seconds / 0 == Duration.Undefined) - assert(Duration.Top / 0 == Duration.Undefined) - assert(Duration.Top / -1 == Duration.Bottom) - assert(Duration.Top / 1 == Duration.Top) - - assert(10.seconds / 2.0 == 5.seconds) - assert(1.seconds / 0.0 == Duration.Top) - assert((-1).seconds / 0.0 == Duration.Bottom) - assert(0.seconds / 0.0 == Duration.Undefined) - assert(Duration.Top / 0.0 == Duration.Undefined) - assert(Duration.Top / -1.0 == Duration.Bottom) - assert(Duration.Top / 1.0 == Duration.Top) - - assert(Duration.Undefined / Double.NaN == Duration.Undefined) - assert(Duration.Top / Double.NaN == Duration.Undefined) - assert(Duration.Bottom / Double.NaN == Duration.Undefined) - } + test("Duration /") { + assert(10.seconds / 2 == 5.seconds) + assert(1.seconds / 0 == Duration.Top) + assert((-1).seconds / 0 == Duration.Bottom) + assert(0.seconds / 0 == Duration.Undefined) + assert(Duration.Top / 0 == Duration.Undefined) + assert(Duration.Top / -1 == Duration.Bottom) + assert(Duration.Top / 1 == Duration.Top) + + assert(10.seconds / 2.0 == 5.seconds) + assert(1.seconds / 0.0 == Duration.Top) + assert((-1).seconds / 0.0 == Duration.Bottom) + assert(0.seconds / 0.0 == Duration.Undefined) + assert(Duration.Top / 0.0 == Duration.Undefined) + assert(Duration.Top / -1.0 == Duration.Bottom) + assert(Duration.Top / 1.0 == Duration.Top) + + assert(Duration.Undefined / Double.NaN == Duration.Undefined) + assert(Duration.Top / Double.NaN == Duration.Undefined) + assert(Duration.Bottom / Double.NaN == Duration.Undefined) + } - "%" in { - assert(10.seconds % 3.seconds == 1.second) - assert(1.second % 300.millis == 100.millis) - assert(1.second % 0.seconds == Duration.Undefined) - assert((-1).second % 0.seconds == Duration.Undefined) - assert(0.seconds % 0.seconds == Duration.Undefined) - assert(Duration.Top % 123.seconds == Duration.Undefined) - assert(Duration.Bottom % 123.seconds == Duration.Undefined) - } + test("Duration %") { + assert(10.seconds % 3.seconds == 1.second) + assert(1.second % 300.millis == 100.millis) + assert(1.second % 0.seconds == Duration.Undefined) + assert((-1).second % 0.seconds == Duration.Undefined) + assert(0.seconds % 0.seconds == Duration.Undefined) + assert(Duration.Top % 123.seconds == Duration.Undefined) + assert(Duration.Bottom % 123.seconds == Duration.Undefined) + } - "unary_-" in { - assert(-((10.seconds).inSeconds) == -10) - assert(-((Long.MinValue + 1).nanoseconds) == Long.MaxValue.nanoseconds) - assert(-(Long.MinValue.nanoseconds) == Duration.Top) - } + test("Duration unary_-") { + assert(-((10.seconds).inSeconds) == -10) + assert(-((Long.MinValue + 1).nanoseconds) == Long.MaxValue.nanoseconds) + assert(-(Long.MinValue.nanoseconds) == Duration.Top) + } - "abs" in { - assert(10.seconds.abs == 10.seconds) - assert((-10).seconds.abs == 10.seconds) - } + test("Duration abs") { + assert(10.seconds.abs == 10.seconds) + assert((-10).seconds.abs == 10.seconds) + } - "afterEpoch" in { - assert(10.seconds.afterEpoch == Time.fromMilliseconds(10000)) - } + test("Duration afterEpoch") { + assert(10.seconds.afterEpoch == Time.fromMilliseconds(10000)) + } - "fromNow" in { - Time.withCurrentTimeFrozen { _ => - assert(10.seconds.fromNow == (Time.now + 10.seconds)) - } - } + test("Duration fromNow") { + Time.withCurrentTimeFrozen { _ => assert(10.seconds.fromNow == (Time.now + 10.seconds)) } + } - "ago" in { - Time.withCurrentTimeFrozen { _ => - assert(10.seconds.ago == (Time.now - 10.seconds)) - } - } + test("Duration ago") { + Time.withCurrentTimeFrozen { _ => assert(10.seconds.ago == (Time.now - 10.seconds)) } + } - "compare" in { - assert(10.seconds < 11.seconds) - assert(10.seconds < 11000.milliseconds) - assert(11.seconds > 10.seconds) - assert(11000.milliseconds > 10.seconds) - assert(10.seconds >= 10.seconds) - assert(10.seconds <= 10.seconds) - assert(new Duration(Long.MaxValue) > 0.seconds) - } + test("Duration compare") { + assert(10.seconds < 11.seconds) + assert(10.seconds < 11000.milliseconds) + assert(11.seconds > 10.seconds) + assert(11000.milliseconds > 10.seconds) + assert(10.seconds >= 10.seconds) + assert(10.seconds <= 10.seconds) + assert(new Duration(Long.MaxValue) > 0.seconds) + } - "equals" in { - assert(Duration.Top == Duration.Top) - assert(Duration.Top != Duration.Bottom) - assert(Duration.Top != Duration.Undefined) + test("Duration equals") { + assert(Duration.Top == Duration.Top) + assert(Duration.Top != Duration.Bottom) + assert(Duration.Top != Duration.Undefined) - assert(Duration.Bottom != Duration.Top) - assert(Duration.Bottom == Duration.Bottom) - assert(Duration.Bottom != Duration.Undefined) + assert(Duration.Bottom != Duration.Top) + assert(Duration.Bottom == Duration.Bottom) + assert(Duration.Bottom != Duration.Undefined) - assert(Duration.Undefined != Duration.Top) - assert(Duration.Undefined != Duration.Bottom) - assert(Duration.Undefined == Duration.Undefined) + assert(Duration.Undefined != Duration.Top) + assert(Duration.Undefined != Duration.Bottom) + assert(Duration.Undefined == Duration.Undefined) - val tenSecs = 10.seconds - assert(tenSecs == tenSecs) - assert(tenSecs == 10.seconds) - assert(tenSecs != 11.seconds) - } + val tenSecs = 10.seconds + assert(tenSecs == tenSecs) + assert(tenSecs == 10.seconds) + assert(tenSecs != 11.seconds) + } - "+ delta" in { - forAll { (a: Long, b: Long) => - val c = BigInt(a) + BigInt(b) - val d = a.nanoseconds + b.nanosecond + test("Duration + delta") { + forAll { (a: Long, b: Long) => + val c = BigInt(a) + BigInt(b) + val d = a.nanoseconds + b.nanosecond - assert( - (c >= Long.MaxValue && d == Duration.Top) || - (c <= Long.MinValue && d == Duration.Bottom) || - (d == c.toLong.nanoseconds) - ) - } + assert( + (c >= Long.MaxValue && d == Duration.Top) || + (c <= Long.MinValue && d == Duration.Bottom) || + (d == c.toLong.nanoseconds) + ) } + } - "- delta" in { - assert(10.seconds - 5.seconds == 5.seconds) - } + test("Duration - delta") { + assert(10.seconds - 5.seconds == 5.seconds) + } - "max" in { - assert((10.seconds max 5.seconds) == 10.seconds) - assert((5.seconds max 10.seconds) == 10.seconds) - } + test("Duration max") { + assert((10.seconds max 5.seconds) == 10.seconds) + assert((5.seconds max 10.seconds) == 10.seconds) + } - "min" in { - assert((10.seconds min 5.seconds) == 5.seconds) - assert((5.seconds min 10.seconds) == 5.seconds) - } + test("Duration min") { + assert((10.seconds min 5.seconds) == 5.seconds) + assert((5.seconds min 10.seconds) == 5.seconds) + } - "moreOrLessEquals" in { - assert(10.seconds.moreOrLessEquals(9.seconds, 1.second) == true) - assert(10.seconds.moreOrLessEquals(11.seconds, 1.second) == true) - assert(10.seconds.moreOrLessEquals(8.seconds, 1.second) == false) - assert(10.seconds.moreOrLessEquals(12.seconds, 1.second) == false) - } + test("Duration moreOrLessEquals") { + assert(10.seconds.moreOrLessEquals(9.seconds, 1.second)) + assert(10.seconds.moreOrLessEquals(11.seconds, 1.second)) + assert(!10.seconds.moreOrLessEquals(8.seconds, 1.second)) + assert(!10.seconds.moreOrLessEquals(12.seconds, 1.second)) + } - "inTimeUnit" in { - assert(23.nanoseconds.inTimeUnit == ((23, TimeUnit.NANOSECONDS))) - assert(23.microseconds.inTimeUnit == ((23000, TimeUnit.NANOSECONDS))) - assert(23.milliseconds.inTimeUnit == ((23, TimeUnit.MILLISECONDS))) - assert(23.seconds.inTimeUnit == ((23, TimeUnit.SECONDS))) - } + test("Duration inTimeUnit") { + assert(23.nanoseconds.inTimeUnit == ((23, TimeUnit.NANOSECONDS))) + assert(23.microseconds.inTimeUnit == ((23000, TimeUnit.NANOSECONDS))) + assert(23.milliseconds.inTimeUnit == ((23, TimeUnit.MILLISECONDS))) + assert(23.seconds.inTimeUnit == ((23, TimeUnit.SECONDS))) + assert(60.seconds.inTimeUnit == ((1, TimeUnit.MINUTES))) + } - "inUnit" in { - assert(23.seconds.inUnit(TimeUnit.SECONDS) == 23L) - assert(23.seconds.inUnit(TimeUnit.MILLISECONDS) == 23000L) - assert(2301.millis.inUnit(TimeUnit.SECONDS) == 2L) - assert(2301.millis.inUnit(TimeUnit.MICROSECONDS) == 2301000L) - assert(4680.nanoseconds.inUnit(TimeUnit.MICROSECONDS) == 4L) - } + test("Duration inUnit") { + assert(23.seconds.inUnit(TimeUnit.SECONDS) == 23L) + assert(23.seconds.inUnit(TimeUnit.MILLISECONDS) == 23000L) + assert(2301.millis.inUnit(TimeUnit.SECONDS) == 2L) + assert(2301.millis.inUnit(TimeUnit.MICROSECONDS) == 2301000L) + assert(4680.nanoseconds.inUnit(TimeUnit.MICROSECONDS) == 4L) + } - "be hashable" in { - val map = new java.util.concurrent.ConcurrentHashMap[Duration, Int] - map.put(23.millis, 23) - assert(map.get(23.millis) == 23) - map.put(44.millis, 44) - assert(map.get(44.millis) == 44) - } + test("Duration be hashable") { + val map = new java.util.concurrent.ConcurrentHashMap[Duration, Int] + map.put(23.millis, 23) + assert(map.get(23.millis) == 23) + map.put(44.millis, 44) + assert(map.get(44.millis) == 44) + } - "toString should display as sums" in { - assert((9999999.seconds).toString == "115.days+17.hours+46.minutes+39.seconds") - } + test("Duration toString should display as sums") { + assert((9999999.seconds).toString == "115.days+17.hours+46.minutes+39.seconds") + } - "toString should handle negative durations" in { - assert((-9999999.seconds).toString == "-115.days-17.hours-46.minutes-39.seconds") - } + test("Duration toString should handle negative durations") { + assert((-9999999.seconds).toString == "-115.days-17.hours-46.minutes-39.seconds") + } - "parse the format from toString" in { - Seq( - -10.minutes, - -9999999.seconds, - 1.day + 3.hours, - 1.day, - 1.nanosecond, - 42.milliseconds, - 9999999.seconds, - Duration.Bottom, - Duration.Top, - Duration.Undefined - ) foreach { d => - assert(Duration.parse(d.toString) == d) - } - } + test("Duration convert java.time.Duration") { + val javaD = java.time.Duration.ofDays(365) + val twitterD = Duration.fromJava(javaD) + assert(twitterD.inDays == 365) - "parse" in { - Seq( - " 1.second" -> 1.second, - "+1.second" -> 1.second, - "-1.second" -> -1.second, - "1.SECOND" -> 1.second, - "1.day - 1.second" -> (1.day - 1.second), - "1.day" -> 1.day, - "1.microsecond" -> 1.microsecond, - "1.millisecond" -> 1.millisecond, - "1.second" -> 1.second, - "1.second+1.minute + 1.day" -> (1.second + 1.minute + 1.day), - "1.second+1.second" -> 2.seconds, - "2.hours" -> 2.hours, - "3.days" -> 3.days, - "321.nanoseconds" -> 321.nanoseconds, - "65.minutes" -> 65.minutes, - "876.milliseconds" -> 876.milliseconds, - "98.seconds" -> 98.seconds, - "Duration.Bottom" -> Duration.Bottom, - "Duration.Top" -> Duration.Top, - "Duration.Undefined" -> Duration.Undefined, - "duration.TOP" -> Duration.Top - ) foreach { - case (s, d) => - assert(Duration.parse(s) == d) - } + assert(Duration.fromJava(javaD) == Duration.fromJava(javaD)) + } + + test("Duration asJavaDuration") { + val t0 = Duration.fromDays(10) + val t1 = t0 + 100.nanosecond + assert(t1.asJava.getNano - t0.asJava.getNano == 100) + } + + test("Duration parse the format from toString") { + Seq( + -10.minutes, + -9999999.seconds, + 1.day + 3.hours, + 1.day, + 1.nanosecond, + 42.milliseconds, + 9999999.seconds, + Duration.Bottom, + Duration.Top, + Duration.Undefined + ).foreach { d => assert(Duration.parse(d.toString) == d) } + } + + test("Duration parse") { + Seq( + " 1.second" -> 1.second, + "+1.second" -> 1.second, + "-1.second" -> -1.second, + "1.SECOND" -> 1.second, + "1.day - 1.second" -> (1.day - 1.second), + "1.day" -> 1.day, + "1.microsecond" -> 1.microsecond, + "1.millisecond" -> 1.millisecond, + "1.second" -> 1.second, + "1.second+1.minute + 1.day" -> (1.second + 1.minute + 1.day), + "1.second+1.second" -> 2.seconds, + "2.hours" -> 2.hours, + "3.days" -> 3.days, + "321.nanoseconds" -> 321.nanoseconds, + "65.minutes" -> 65.minutes, + "876.milliseconds" -> 876.milliseconds, + "98.seconds" -> 98.seconds, + "Duration.Bottom" -> Duration.Bottom, + "Duration.Top" -> Duration.Top, + "Duration.Undefined" -> Duration.Undefined, + "duration.TOP" -> Duration.Top + ).foreach { + case (s, d) => + assert(Duration.parse(s) == d) } + } - "reject obvious human impostors" in { - intercept[NumberFormatException] { - Seq( - "", - "++1.second", - "1. second", - "1.milli", - "1.s", - "10.stardates", - "2.minutes 1.second", - "98 milliseconds", - "98 millisecons", - "99.minutes +" - ) foreach { s => - Duration.parse(s) - } + test("Duration reject obvious human impostors") { + Seq( + "", + "++1.second", + "1. second", + "1.milli", + "1.s", + "10.stardates", + "2.minutes 1.second", + "98 milliseconds", + "98 millisecons", + "99.minutes +" + ).foreach { s => + assertThrows[NumberFormatException] { + Duration.parse(s) } } } - "Top" should { - "Be scaling resistant" in { - assert(Duration.Top / 100 == Duration.Top) - assert(Duration.Top * 100 == Duration.Top) - } + test("Top should be scaling resistant") { + assert(Duration.Top / 100 == Duration.Top) + assert(Duration.Top * 100 == Duration.Top) + } - "-Top == Bottom" in { - assert(-Duration.Top == Duration.Bottom) - } + test("-Top == Bottom") { + assert(-Duration.Top == Duration.Bottom) + } - "--Top == Top" in { - assert(-(-Duration.Top) == Duration.Top) - } + test("--Top == Top") { + assert(-(-Duration.Top) == Duration.Top) } } diff --git a/util-core/src/test/scala/com/twitter/util/EventTest.scala b/util-core/src/test/scala/com/twitter/util/EventTest.scala index 416d336cf1..da9bc69077 100644 --- a/util-core/src/test/scala/com/twitter/util/EventTest.scala +++ b/util-core/src/test/scala/com/twitter/util/EventTest.scala @@ -1,14 +1,11 @@ package com.twitter.util import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} -import java.util.concurrent.Executors -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.junit.runner.RunWith +import java.util.concurrent.{CountDownLatch, Executors} import scala.collection.mutable +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class EventTest extends FunSuite { +class EventTest extends AnyFunSuite { test("pub/sub while active") { val e = Event[Int]() @@ -138,15 +135,11 @@ class EventTest extends FunSuite { def register(w: Witness[Int]) = { n += 1 w.notify(1) - Closable.make { _ => - n -= 1; Future.Done - } + Closable.make { _ => n -= 1; Future.Done } } } - val e12 = e1 mergeMap { _ => - e2 - } + val e12 = e1 mergeMap { _ => e2 } val ref = new AtomicReference(Seq.empty[Int]) val closable = e12.build.register(Witness(ref)) @@ -280,9 +273,7 @@ class EventTest extends FunSuite { } def ite[T](i: Var[Boolean], t: Var[T], e: Var[T]) = - i flatMap { i => - if (i) t else e - } + i flatMap { i => if (i) t else e } val b = Var(true) val x = Var(7) @@ -330,11 +321,8 @@ class EventTest extends FunSuite { test("Event.dedupWith") { val e = Event[Int]() - val ref = new AtomicReference[IndexedSeq[Int]] - - e.dedupWith { (a, b) => - a >= b - } + val ref = new AtomicReference[Seq[Int]] + e.dedupWith { (a, b) => a >= b } .build .register(Witness(ref)) e.notify(0) @@ -350,9 +338,8 @@ class EventTest extends FunSuite { test("Event.dedup") { val e = Event[Int]() - val ref = new AtomicReference[IndexedSeq[Int]] - - e.dedup.build.register(Witness(ref)) + val ref = new AtomicReference[Seq[Int]] + e.dedup.buildAny[IndexedSeq[Int]].register(Witness(ref)) e.notify(0) e.notify(0) e.notify(1) @@ -362,4 +349,5 @@ class EventTest extends FunSuite { assert(ref.get() == List(0, 1, 0)) } + } diff --git a/util-core/src/test/scala/com/twitter/util/ExecutorWorkQueueFiberTest.scala b/util-core/src/test/scala/com/twitter/util/ExecutorWorkQueueFiberTest.scala new file mode 100644 index 0000000000..3a1b15ee5d --- /dev/null +++ b/util-core/src/test/scala/com/twitter/util/ExecutorWorkQueueFiberTest.scala @@ -0,0 +1,70 @@ +package com.twitter.util + +import java.util.concurrent.CountDownLatch +import java.util.concurrent.Executor +import java.util.concurrent.ExecutorService +import java.util.concurrent.Executors +import java.util.concurrent.atomic.AtomicInteger +import org.scalatest.funsuite.AnyFunSuite + +class ExecutorWorkQueueFiberTest extends AnyFunSuite { + + private class NoOpMetrics extends WorkQueueFiber.FiberMetrics { + def flushIncrement(): Unit = () + def fiberCreated(): Unit = () + def threadLocalSubmitIncrement(): Unit = () + def taskSubmissionIncrement(): Unit = () + def schedulerSubmissionIncrement(): Unit = () + } + + private class TestPool(exec: ExecutorService) extends Executor { + val poolSubmitCount = new AtomicInteger(0) + override def execute(command: Runnable): Unit = { + poolSubmitCount.incrementAndGet() + exec.execute(command) + } + + def shutdown() = exec.shutdown() + } + + test("fiber submits tasks to underlying future pool in single runnable") { + val pool = new TestPool(Executors.newFixedThreadPool(1)) + val fiber = new ExecutorWorkQueueFiber(pool, new NoOpMetrics) + + val latch = new CountDownLatch(1) + val N = 5 + for (_ <- 0 until N) { + fiber.submitTask { () => + // wait to finish task execution so that all task submissions queue up in + // the fiber and are submitted to the pool as a single runnable + latch.await() + } + } + latch.countDown() + assert(pool.poolSubmitCount.get() == 1) + + pool.shutdown() + } + + test("fiber submits tasks individually to underlying future pool") { + val pool = new TestPool(Executors.newFixedThreadPool(1)) + val fiber = new ExecutorWorkQueueFiber(pool, new NoOpMetrics) + + val N = 5 + for (_ <- 0 until N) { + val latch = new CountDownLatch(1) + fiber.submitTask { () => + latch.countDown() + } + // wait for the task to finish executing so that the next task gets its own + // submission to the pool + latch.await() + // give the executing thread some time to finish up resource tracking and fully + // exit out of its work loop + Thread.sleep(1) + } + assert(pool.poolSubmitCount.get() == N) + + pool.shutdown() + } +} diff --git a/util-core/src/test/scala/com/twitter/util/FutureOfferTest.scala b/util-core/src/test/scala/com/twitter/util/FutureOfferTest.scala index 2c4f7eb8de..04efa1ad32 100644 --- a/util-core/src/test/scala/com/twitter/util/FutureOfferTest.scala +++ b/util-core/src/test/scala/com/twitter/util/FutureOfferTest.scala @@ -1,28 +1,23 @@ package com.twitter.concurrent -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import com.twitter.util.Promise +import com.twitter.util.Return +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.mockito.MockitoSugar -import com.twitter.util.{Promise, Return} - -@RunWith(classOf[JUnitRunner]) -class FutureOfferTest extends WordSpec with MockitoSugar { - "Future.toOffer" should { - "activate when future is satisfied (poll)" in { - val p = new Promise[Int] - val o = p.toOffer - assert(o.prepare().poll == None) - p() = Return(123) - assert(o.prepare().poll match { - case Some(Return(tx)) => - tx.ack().poll match { - case Some(Return(Tx.Commit(Return(123)))) => true - case _ => false - } - case _ => false - }) - } +class FutureOfferTest extends AnyFunSuite with MockitoSugar { + test("Future.toOffer should activate when future is satisfied (poll)") { + val p = new Promise[Int] + val o = p.toOffer + assert(o.prepare().poll == None) + p() = Return(123) + assert(o.prepare().poll match { + case Some(Return(tx)) => + tx.ack().poll match { + case Some(Return(Tx.Commit(Return(123)))) => true + case _ => false + } + case _ => false + }) } } diff --git a/util-core/src/test/scala/com/twitter/util/FuturePoolTest.scala b/util-core/src/test/scala/com/twitter/util/FuturePoolTest.scala index a49f81fdd2..6c13d0712d 100644 --- a/util-core/src/test/scala/com/twitter/util/FuturePoolTest.scala +++ b/util-core/src/test/scala/com/twitter/util/FuturePoolTest.scala @@ -1,21 +1,51 @@ package com.twitter.util -import com.twitter.conversions.time._ +import com.twitter.concurrent.Scheduler +import com.twitter.concurrent.ForkingScheduler +import com.twitter.conversions.DurationOps._ import java.util.concurrent.{Future => JFuture, _} -import org.junit.runner.RunWith -import org.scalatest.FunSuite import org.scalatest.concurrent.Eventually -import org.scalatest.junit.JUnitRunner -import org.scalatest.time.{Millis, Seconds, Span} +import org.scalatest.time.Millis +import org.scalatest.time.Seconds +import org.scalatest.time.Span import scala.runtime.NonLocalReturnControl import scala.util.control.NonFatal +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class FuturePoolTest extends FunSuite with Eventually { +class FuturePoolTest extends AnyFunSuite with Eventually { - implicit override val patienceConfig = + implicit override val patienceConfig: PatienceConfig = PatienceConfig(timeout = scaled(Span(15, Seconds)), interval = scaled(Span(5, Millis))) + class MockTracker extends ResourceTracker { + + @volatile var updates: Long = 0 + + override def addCpuTime(time: Long): Unit = { + updates = updates + 1 + } + } + + test("ResourceTracker, when set -- FuturePool tracks usage on execution") { + if (ResourceTracker.threadCpuTimeSupported) { + val executor = Executors.newFixedThreadPool(1) + val pool = FuturePool(executor) + + val mockTracker = new MockTracker + val latch = new CountDownLatch(1) + val result = ResourceTracker.let(mockTracker) { + pool { + // Perform a (potentially) blocking call. + latch.await() + } + } + + latch.countDown() + Await.result(result) + assert(mockTracker.updates == 1) + } + } + test("FuturePool should dispatch to another thread") { val executor = Executors.newFixedThreadPool(1) val pool = FuturePool(executor) @@ -58,7 +88,8 @@ class FuturePoolTest extends FunSuite with Eventually { val result1 = pool { runCount.incrementAndGet(); Await.result(source1) } val result2 = pool { runCount.incrementAndGet(); Await.result(source2) } - result2.raise(new Exception) + val exc = new CancellationException + result2.raise(exc) source1.setValue(1) // The executor will run the task for result 2, but the wrapper @@ -93,7 +124,7 @@ class FuturePoolTest extends FunSuite with Eventually { runCount.get } - startedLatch.await(1.second) + startedLatch.await(1, TimeUnit.SECONDS) result.raise(new Exception) cancelledLatch.countDown() @@ -173,10 +204,10 @@ class FuturePoolTest extends FunSuite with Eventually { } class PoolCtx { - val executor = Executors.newFixedThreadPool(1) - val pool = FuturePool(executor) + val executor: ExecutorService = Executors.newFixedThreadPool(1) + val pool: ExecutorServiceFuturePool = FuturePool(executor) - val pools = Seq(FuturePool.immediatePool, pool) + val pools: Seq[FuturePool] = Seq(FuturePool.immediatePool, pool) } test("handles NonLocalReturnControl properly") { @@ -202,32 +233,148 @@ class FuturePoolTest extends FunSuite with Eventually { // But we want to make sure it will have the correct behavior. // We compromise by roughly creating an ExecutorService/FuturePool // that behaves the same. - val executor = Executors.newCachedThreadPool() + val executor = Executors.newFixedThreadPool(1) val pool = new ExecutorServiceFuturePool(executor) // verify the initial state assert(pool.poolSize == 0) assert(pool.numActiveTasks == 0) assert(pool.numCompletedTasks == 0) + assert(pool.numPendingTasks == 0) // execute a task we can control val latch = new CountDownLatch(1) val future = pool { - latch.await(10.seconds) + latch.await(10, TimeUnit.SECONDS) true } + pool(()) // pending task + eventually { assert(pool.poolSize == 1) } eventually { assert(pool.numActiveTasks == 1) } eventually { assert(pool.numCompletedTasks == 0) } + eventually { assert(pool.numPendingTasks == 1) } // let the task complete latch.countDown() Await.ready(future, 5.seconds) eventually { assert(pool.poolSize == 1) } eventually { assert(pool.numActiveTasks == 0) } - eventually { assert(pool.numCompletedTasks == 1) } + eventually { assert(pool.numCompletedTasks == 2) } + eventually { assert(pool.numPendingTasks == 0) } // cleanup. executor.shutdown() } + class TestForkingScheduler(redirect: Boolean) extends ForkingScheduler { + var concurrency = Option.empty[Int] + var maxWaiters = Option.empty[Int] + var forked = Option.empty[Future[_]] + override def fork[T](f: => Future[T]) = { + val r = f + forked = Some(r) + r + } + override def tryFork[T](f: => Future[T]): Future[Option[T]] = ??? + override def redirectFuturePools(): Boolean = redirect + override def blocking[T](f: => T)(implicit perm: Awaitable.CanAwait): T = f + override def flush(): Unit = {} + override def withMaxSyncConcurrency(v: Int, w: Int): ForkingScheduler = { + concurrency = Some(v) + maxWaiters = Some(w) + this + } + + override def submit(r: Runnable): Unit = r.run() + + override def fork[T](executor: Executor)(f: => Future[T]) = ??? + override def asExecutorService(): ExecutorService = ??? + override def numDispatches: Long = ??? + } + + test( + "FuturePool should be redirected to the forking scheduler if specified by the scheduler - unbounded queue") { + val concurrency = 1 + val scheduler = new TestForkingScheduler(redirect = true) + val prevScheduler = Scheduler() + val executor = Executors.newFixedThreadPool(concurrency) + Scheduler.setUnsafe(scheduler) + try { + val pool = FuturePool(executor) + var executed = false + val result = + pool { + executed = true + 1 + } + assert(Await.result(result, 1.second) == 1) + assert(Await.result(scheduler.forked.get, 1.second) == 1) + assert(scheduler.concurrency == Some(concurrency)) + assert(scheduler.maxWaiters == Some(Int.MaxValue)) + assert(executed) + } finally { + Scheduler.setUnsafe(prevScheduler) + executor.shutdown() + } + } + + test( + "FuturePool should be redirected to the forking scheduler if specified by the scheduler - bounded queue") { + val concurrency = 1 + val maxWaiters = 10 + val scheduler = new TestForkingScheduler(redirect = true) + val prevScheduler = Scheduler() + val executor = + new ThreadPoolExecutor( + concurrency, + concurrency, + 0L, + TimeUnit.MILLISECONDS, + new LinkedBlockingQueue(maxWaiters)) + Scheduler.setUnsafe(scheduler) + try { + val pool = FuturePool(executor) + var executed = false + val result = + pool { + executed = true + 1 + } + assert(Await.result(result, 1.second) == 1) + assert(Await.result(scheduler.forked.get, 1.second) == 1) + assert(scheduler.concurrency == Some(concurrency)) + assert(scheduler.maxWaiters == Some(maxWaiters)) + assert(executed) + } finally { + Scheduler.setUnsafe(prevScheduler) + executor.shutdown() + } + } + + test( + "FuturePool should not be redirected to the forking scheduler if specified by the scheduler") { + val concurrency = 1 + val scheduler = new TestForkingScheduler(redirect = false) + val executor = Executors.newFixedThreadPool(concurrency) + val prevScheduler = Scheduler() + Scheduler.setUnsafe(scheduler) + try { + val pool = FuturePool(executor) + var executed = false + val result = + pool { + executed = true + 1 + } + assert(Await.result(result, 1.second) == 1) + assert(scheduler.forked.isEmpty) + assert(scheduler.concurrency.isEmpty) + assert(scheduler.maxWaiters.isEmpty) + assert(executed) + } finally { + Scheduler.setUnsafe(prevScheduler) + executor.shutdown() + } + } + } diff --git a/util-core/src/test/scala/com/twitter/util/FutureTest.scala b/util-core/src/test/scala/com/twitter/util/FutureTest.scala index 78762eabcf..387abb7b1d 100644 --- a/util-core/src/test/scala/com/twitter/util/FutureTest.scala +++ b/util-core/src/test/scala/com/twitter/util/FutureTest.scala @@ -1,17 +1,23 @@ package com.twitter.util -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import java.util.concurrent.ConcurrentLinkedQueue +import java.util.concurrent.ExecutionException +import java.util.concurrent.{Future => JFuture} import java.util.concurrent.atomic.AtomicInteger -import org.mockito.Matchers.any -import org.mockito.Mockito.{never, times, verify, when} +import org.mockito.ArgumentMatchers.any +import org.mockito.Mockito.never +import org.mockito.Mockito.times +import org.mockito.Mockito.verify +import org.mockito.Mockito.when import org.mockito.invocation.InvocationOnMock import org.mockito.stubbing.Answer -import org.scalacheck.{Arbitrary, Gen} -import org.scalatest.WordSpec -import org.scalatest.mockito.MockitoSugar -import org.scalatest.prop.GeneratorDrivenPropertyChecks -import scala.collection.JavaConverters._ +import org.scalacheck.Arbitrary.arbitrary +import org.scalacheck.Gen +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.mockito.MockitoSugar +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import scala.jdk.CollectionConverters._ import scala.runtime.NonLocalReturnControl import scala.util.Random import scala.util.control.ControlThrowable @@ -38,7 +44,7 @@ private object FutureTest { } } -class FutureTest extends WordSpec with MockitoSugar with GeneratorDrivenPropertyChecks { +class FutureTest extends AnyFunSuite with MockitoSugar with ScalaCheckDrivenPropertyChecks { import FutureTest._ implicit class FutureMatcher[A](future: Future[A]) { @@ -71,6 +77,75 @@ class FutureTest extends WordSpec with MockitoSugar with GeneratorDrivenProperty } } + trait TimesHelper { + val queue = new ConcurrentLinkedQueue[Promise[Unit]] + var complete = false + var failure = false + var ninterrupt = 0 + val iteration: Future[Unit] = Future.times(3) { + val promise = new Promise[Unit] + promise.setInterruptHandler { case _ => ninterrupt += 1 } + queue add promise + promise + } + iteration + .onSuccess { _ => complete = true } + .onFailure { _ => failure = true } + assert(!complete) + assert(!failure) + } + + trait WhileDoHelper { + var i = 0 + val queue = new ConcurrentLinkedQueue[HandledPromise[Unit]] + var complete = false + var failure = false + val iteration: Future[Unit] = Future.whileDo(i < 3) { + i += 1 + val promise = new HandledPromise[Unit] + queue.add(promise) + promise + } + + iteration + .onSuccess { _ => complete = true } + .onFailure { _ => failure = true } + assert(!complete) + assert(!failure) + } + + class TraverseTestSpy() { + var goWasCalled = false + var promise: Promise[Int] = Promise[Int]() + val go: () => Promise[Int] = () => { + goWasCalled = true + promise + } + } + + trait CollectHelper { + val p0, p1 = new HandledPromise[Int] + val f: Future[Seq[Int]] = Future.collect(Seq(p0, p1)) + assert(!f.isDefined) + } + + trait CollectToTryHelper { + val p0, p1 = new HandledPromise[Int] + val f: Future[Seq[Try[Int]]] = Future.collectToTry(Seq(p0, p1)) + assert(!f.isDefined) + } + + trait SelectHelper { + var nhandled = 0 + val p0, p1 = new HandledPromise[Int] + val f: Future[Int] = p0.select(p1) + assert(!f.isDefined) + } + + trait PollHelper { + val p = new Promise[Int] + } + class FatalException extends ControlThrowable trait MkConst { @@ -79,1683 +154,1641 @@ class FutureTest extends WordSpec with MockitoSugar with GeneratorDrivenProperty def exception[A](exc: Throwable): Future[A] = this(Throw(exc)) } - def test(name: String, const: MkConst): Unit = { - s"object Future ($name)" when { - "times" should { - trait TimesHelper { - val queue = new ConcurrentLinkedQueue[Promise[Unit]] - var complete = false - var failure = false - var ninterrupt = 0 - val iteration: Future[Unit] = Future.times(3) { - val promise = new Promise[Unit] - promise.setInterruptHandler { case _ => ninterrupt += 1 } - queue add promise - promise - } - iteration.onSuccess { _ => - complete = true - }.onFailure { _ => - failure = true - } - assert(!complete) - assert(!failure) - } + trait MonitoredHelper { + val inner = new HandledPromise[Int] + val exc = new Exception("some exception") + } - "when everything succeeds" in { - new TimesHelper { - queue.poll().setDone() - assert(!complete) - assert(!failure) + def tests(name: String, const: MkConst): Unit = { + test(s"object Future ($name) times when everything succeeds") { + new TimesHelper { + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setDone() - assert(!complete) - assert(!failure) + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setDone() - assert(complete) - assert(!failure) - } - } + queue.poll().setDone() + assert(complete) + assert(!failure) + } + } - "when some succeed and some fail" in { - new TimesHelper { - queue.poll().setDone() - assert(!complete) - assert(!failure) + test(s"object Future ($name) times when some succeed and some fail") { + new TimesHelper { + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setException(new Exception("")) - assert(!complete) - assert(failure) - } - } + queue.poll().setException(new Exception("")) + assert(!complete) + assert(failure) + } + } - "when interrupted" in { - new TimesHelper { - assert(ninterrupt == 0) - iteration.raise(new Exception) - for (i <- 1 to 3) { - assert(ninterrupt == i) - queue.poll().setDone() - } - } + test(s"object Future ($name) times when interrupted") { + new TimesHelper { + assert(ninterrupt == 0) + iteration.raise(new Exception) + for (i <- 1 to 3) { + assert(ninterrupt == i) + queue.poll().setDone() } } + } - "when" in { - var i = 0 - - await { - Future.when(false) { - Future { i += 1 } - } - } - assert(i == 0) + test(s"object Future ($name) when") { + var i = 0 - await { - Future.when(true) { - Future { i += 1 } - } + await { + Future.when(false) { + Future { i += 1 } } - assert(i == 1) - } - - "whileDo" should { - trait WhileDoHelper { - var i = 0 - val queue = new ConcurrentLinkedQueue[HandledPromise[Unit]] - var complete = false - var failure = false - val iteration: Future[Unit] = Future.whileDo(i < 3) { - i += 1 - val promise = new HandledPromise[Unit] - queue.add(promise) - promise - } + } + assert(i == 0) - iteration.onSuccess { _ => - complete = true - }.onFailure { _ => - failure = true - } - assert(!complete) - assert(!failure) + await { + Future.when(true) { + Future { i += 1 } } + } + assert(i == 1) + } - "when everything succeeds" in { - new WhileDoHelper { - queue.poll().setDone() - assert(!complete) - assert(!failure) + test(s"object Future ($name) whileDo when everything succeeds") { + new WhileDoHelper { + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setDone() - assert(!complete) - assert(!failure) + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setDone() + queue.poll().setDone() - assert(complete) - assert(!failure) - } - } + assert(complete) + assert(!failure) + } + } - "when some succeed and some fail" in { - new WhileDoHelper { - queue.poll().setDone() - assert(!complete) - assert(!failure) + test(s"object Future ($name) whileDo when some succeed and some fail") { + new WhileDoHelper { + queue.poll().setDone() + assert(!complete) + assert(!failure) - queue.poll().setException(new Exception("")) - assert(!complete) - assert(failure) - } - } + queue.poll().setException(new Exception("")) + assert(!complete) + assert(failure) + } + } - "when interrupted" in { - new WhileDoHelper { - assert(!queue.asScala.exists(_.handled.isDefined)) - iteration.raise(new Exception) - assert(queue.asScala.forall(_.handled.isDefined)) - } - } + test(s"object Future ($name) whileDo when interrupted") { + new WhileDoHelper { + assert(!queue.asScala.exists(_.handled.isDefined)) + iteration.raise(new Exception) + assert(queue.asScala.forall(_.handled.isDefined)) } + } - "proxyTo" should { - "reject satisfied promises" in { - val str = "um um excuse me um" - val p1 = new Promise[String]() - p1.update(Return(str)) + test(s"object Future ($name) proxyTo should reject satisfied promises") { + val str = "um um excuse me um" + val p1 = new Promise[String]() + p1.update(Return(str)) - val p2 = new Promise[String]() - val ex = intercept[IllegalStateException] { p2.proxyTo(p1) } - assert(ex.getMessage.contains(str)) - } + val p2 = new Promise[String]() + val ex = intercept[IllegalStateException] { p2.proxyTo(p1) } + assert(ex.getMessage.contains(str)) + } - "proxies success" in { - val p1 = new Promise[Int]() - val p2 = new Promise[Int]() - p2.proxyTo(p1) - p2.update(Return(5)) - assert(5 == await(p1)) - assert(5 == await(p2)) - } + test(s"object Future ($name) proxyTo proxies success") { + val p1 = new Promise[Int]() + val p2 = new Promise[Int]() + p2.proxyTo(p1) + p2.update(Return(5)) + assert(5 == await(p1)) + assert(5 == await(p2)) + } - "proxies failure" in { - val p1 = new Promise[Int]() - val p2 = new Promise[Int]() - p2.proxyTo(p1) - - val t = new RuntimeException("wurmp") - p2.update(Throw(t)) - val ex1 = intercept[RuntimeException] { await(p1) } - assert(ex1.getMessage == t.getMessage) - val ex2 = intercept[RuntimeException] { await(p2) } - assert(ex2.getMessage == t.getMessage) - } + test(s"object Future ($name) proxyTo proxies failure") { + val p1 = new Promise[Int]() + val p2 = new Promise[Int]() + p2.proxyTo(p1) + + val t = new RuntimeException("wurmp") + p2.update(Throw(t)) + val ex1 = intercept[RuntimeException] { await(p1) } + assert(ex1.getMessage == t.getMessage) + val ex2 = intercept[RuntimeException] { await(p2) } + assert(ex2.getMessage == t.getMessage) + } + + def batchedTests() = { + implicit val batchedTimer: MockTimer = new MockTimer + val batchedResult = Seq(4, 5, 6) + + test(s"object Future ($name) batched should execute after threshold is reached") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + + val batcher = Future.batched(3)(f) + + when(m.apply(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) + batcher(1) + verify(m, never()).apply(any[Seq[Int]]) + batcher(2) + verify(m, never()).apply(any[Seq[Int]]) + batcher(3) + verify(m).apply(Iterable(1, 2, 3)) } - "batched" should { - implicit val timer: MockTimer = new MockTimer - val result = Seq(4, 5, 6) + test( + s"object Future ($name) batched should execute after bufSizeFraction threshold is reached") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + + val batcher = Future.batched(3, sizePercentile = 0.67f)(f) + when(m.apply(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) + batcher(1) + verify(m, never()).apply(any[Seq[Int]]) + batcher(2) + verify(m).apply(Iterable(1, 2)) + } - "execute after threshold is reached" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3)(f) + test(s"object Future ($name) batched should treat bufSizeFraction return value < 0.0f as 1") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - when(f.apply(Seq(1, 2, 3))).thenReturn(Future.value(result)) - batcher(1) - verify(f, never()).apply(any[Seq[Int]]) - batcher(2) - verify(f, never()).apply(any[Seq[Int]]) - batcher(3) - verify(f).apply(Seq(1, 2, 3)) - } + val batcher = Future.batched(3, sizePercentile = 0.4f)(f) - "execute after bufSizeFraction threshold is reached" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3, sizePercentile = 0.67f)(f) + when(m.apply(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) + batcher(1) + verify(m).apply(Iterable(1)) + } - when(f.apply(Seq(1, 2, 3))).thenReturn(Future.value(result)) - batcher(1) - verify(f, never()).apply(any[Seq[Int]]) - batcher(2) - verify(f).apply(Seq(1, 2)) - } + test( + s"object Future ($name) batched should treat bufSizeFraction return value > 1.0f should return maxSizeThreshold") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + + val batcher = Future.batched(3, sizePercentile = 1.3f)(f) + + when(m.apply(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) + batcher(1) + verify(m, never()).apply(any[Seq[Int]]) + batcher(2) + verify(m, never()).apply(any[Seq[Int]]) + batcher(3) + verify(m).apply(Iterable(1, 2, 3)) + } - "treat bufSizeFraction return value < 0.0f as 1" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3, sizePercentile = 0.4f)(f) + test(s"object Future ($name) batched should execute after time threshold") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - when(f.apply(Seq(1, 2, 3))).thenReturn(Future.value(result)) + val batcher = Future.batched(3, 3.seconds)(f) + + Time.withCurrentTimeFrozen { control => + when(m(Iterable(1))).thenReturn(Future.value(Seq(4))) batcher(1) - verify(f).apply(Seq(1)) + verify(m, never()).apply(any[Seq[Int]]) + + control.advance(1.second) + batchedTimer.tick() + verify(m, never()).apply(any[Seq[Int]]) + + control.advance(1.second) + batchedTimer.tick() + verify(m, never()).apply(any[Seq[Int]]) + + control.advance(1.second) + batchedTimer.tick() + verify(m).apply(Iterable(1)) } + } + + test(s"object Future ($name) batched should only execute once if both are reached") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - "treat bufSizeFraction return value > 1.0f should return maxSizeThreshold" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3, sizePercentile = 1.3f)(f) + val batcher = Future.batched(3)(f) - when(f.apply(Seq(1, 2, 3))).thenReturn(Future.value(result)) + Time.withCurrentTimeFrozen { control => + when(m(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) batcher(1) - verify(f, never()).apply(any[Seq[Int]]) batcher(2) - verify(f, never()).apply(any[Seq[Int]]) batcher(3) - verify(f).apply(Seq(1, 2, 3)) + control.advance(10.seconds) + batchedTimer.tick() + + verify(m).apply(Iterable(1, 2, 3)) } + } - "execute after time threshold" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3, 3.seconds)(f) + test(s"object Future ($name) batched should execute when flushBatch is called") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - Time.withCurrentTimeFrozen { control => - when(f(Seq(1))).thenReturn(Future.value(Seq(4))) - batcher(1) - verify(f, never()).apply(any[Seq[Int]]) + val batcher = Future.batched(4)(f) - control.advance(1.second) - timer.tick() - verify(f, never()).apply(any[Seq[Int]]) + batcher(1) + batcher(2) + batcher(3) + batcher.flushBatch() - control.advance(1.second) - timer.tick() - verify(f, never()).apply(any[Seq[Int]]) + verify(m).apply(Iterable(1, 2, 3)) + } - control.advance(1.second) - timer.tick() - verify(f).apply(Seq(1)) - } - } + test( + s"object Future ($name) batched should only execute for remaining items when flushBatch is called after size threshold is reached") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - "only execute once if both are reached" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3)(f) + val batcher = Future.batched(4)(f) - Time.withCurrentTimeFrozen { control => - when(f(Seq(1, 2, 3))).thenReturn(Future.value(result)) - batcher(1) - batcher(2) - batcher(3) - control.advance(10.seconds) - timer.tick() + batcher(1) + batcher(2) + batcher(3) + batcher(4) + batcher(5) + verify(m, times(1)).apply(Iterable(1, 2, 3, 4)) - verify(f).apply(Seq(1, 2, 3)) - } - } + batcher.flushBatch() + verify(m, times(1)).apply(Iterable(5)) + } - "execute when flushBatch is called" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(4)(f) + test( + s"object Future ($name) batched should only execute once when time threshold is reached after flushBatch is called") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + val batcher = Future.batched(4, 3.seconds)(f) + + Time.withCurrentTimeFrozen { control => batcher(1) batcher(2) batcher(3) batcher.flushBatch() + control.advance(10.seconds) + batchedTimer.tick() - verify(f).apply(Seq(1, 2, 3)) + verify(m, times(1)).apply(Iterable(1, 2, 3)) } + } - "only execute for remaining items when flushBatch is called after size threshold is reached" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(4)(f) + test( + s"object Future ($name) batched should only execute once when time threshold is reached before flushBatch is called") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + val batcher = Future.batched(4, 3.seconds)(f) + + Time.withCurrentTimeFrozen { control => batcher(1) batcher(2) batcher(3) - batcher(4) - batcher(5) - verify(f, times(1)).apply(Seq(1, 2, 3, 4)) - + control.advance(10.seconds) + batchedTimer.tick() batcher.flushBatch() - verify(f, times(1)).apply(Seq(5)) - } - - "only execute once when time threshold is reached after flushBatch is called" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(4, 3.seconds)(f) - Time.withCurrentTimeFrozen { control => - batcher(1) - batcher(2) - batcher(3) - batcher.flushBatch() - control.advance(10.seconds) - timer.tick() - - verify(f, times(1)).apply(Seq(1, 2, 3)) - } + verify(m, times(1)).apply(Iterable(1, 2, 3)) } + } - "only execute once when time threshold is reached before flushBatch is called" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(4, 3.seconds)(f) + test(s"object Future ($name) batched should propagates results") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - Time.withCurrentTimeFrozen { control => - batcher(1) - batcher(2) - batcher(3) - control.advance(10.seconds) - timer.tick() - batcher.flushBatch() + val batcher = Future.batched(3)(f) - verify(f, times(1)).apply(Seq(1, 2, 3)) - } - } + Time.withCurrentTimeFrozen { _ => + when(m(Iterable(1, 2, 3))).thenReturn(Future.value(batchedResult)) + val res1 = batcher(1) + assert(!res1.isDefined) + val res2 = batcher(2) + assert(!res2.isDefined) + val res3 = batcher(3) + assert(res1.isDefined) + assert(res2.isDefined) + assert(res3.isDefined) - "propagates results" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3)(f) + assert(await(res1) == 4) + assert(await(res2) == 5) + assert(await(res3) == 6) - Time.withCurrentTimeFrozen { _ => - when(f(Seq(1, 2, 3))).thenReturn(Future.value(result)) - val res1 = batcher(1) - assert(!res1.isDefined) - val res2 = batcher(2) - assert(!res2.isDefined) - val res3 = batcher(3) - assert(res1.isDefined) - assert(res2.isDefined) - assert(res3.isDefined) - - assert(await(res1) == 4) - assert(await(res2) == 5) - assert(await(res3) == 6) - - verify(f).apply(Seq(1, 2, 3)) - } + verify(m).apply(Iterable(1, 2, 3)) } + } - "not block other batches" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3)(f) - - Time.withCurrentTimeFrozen { _ => - val blocker = new Promise[Unit] - val thread = new Thread { - override def run(): Unit = { - when(f(result)).thenReturn(Future.value(Seq(7, 8, 9))) - batcher(4) - batcher(5) - batcher(6) - verify(f).apply(result) - blocker.setValue(()) - } - } - - when(f(Seq(1, 2, 3))).thenAnswer { - new Answer[Future[Seq[Int]]] { - def answer(invocation: InvocationOnMock): Future[Seq[Int]] = { - thread.start() - await(blocker) - Future.value(result) - } - } + test(s"object Future ($name) batched should not block other batches") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) + + val batcher = Future.batched(3)(f) + + Time.withCurrentTimeFrozen { _ => + val blocker = new Promise[Unit] + val thread = new Thread { + override def run(): Unit = { + when(m(batchedResult)).thenReturn(Future.value(Seq(7, 8, 9))) + batcher(4) + batcher(5) + batcher(6) + verify(m).apply(batchedResult) + blocker.setValue(()) } - - batcher(1) - batcher(2) - batcher(3) - verify(f).apply(Seq(1, 2, 3)) } - } - "swallow exceptions" in { - val f = mock[Seq[Int] => Future[Seq[Int]]] - val batcher = Future.batched(3)(f) - - when(f(Seq(1, 2, 3))).thenAnswer { - new Answer[Unit] { - def answer(invocation: InvocationOnMock): Unit = { - throw new Exception + when(m(Seq(1, 2, 3))).thenAnswer { + new Answer[Future[Seq[Int]]] { + def answer(invocation: InvocationOnMock): Future[Seq[Int]] = { + thread.start() + await(blocker) + Future.value(batchedResult) } } } batcher(1) batcher(2) - batcher(3) // Success here implies no exception was thrown. + batcher(3) + verify(m).apply(Iterable(1, 2, 3)) } } - "interruptible" should { - "properly ignore the underlying future on interruption" in { - val p = Promise[Unit]() - val i = p.interruptible() - val e = new Exception - i.raise(e) - p.setDone() - assert(p.poll.contains(Return(()))) - assert(i.poll.contains(Throw(e))) - } - - "respect the underlying future" in { - val p = Promise[Unit]() - val i = p.interruptible() - p.setDone() - assert(p.poll.contains(Return(()))) - assert(i.poll.contains(Return(()))) - } + test(s"object Future ($name) batched should swallow exceptions") { + val m = mock[Any => Future[Seq[Int]]] + val f = mock[scala.collection.Seq[Int] => Future[Seq[Int]]] + when(f.compose(any())).thenReturn(m) - "do nothing for const" in { - val f = const.value(()) - val i = f.interruptible() - i.raise(new Exception()) - assert(f.poll.contains(Return(()))) - assert(i.poll.contains(Return(()))) - } - } + val batcher = Future.batched(3)(f) - "traverseSequentially" should { - class TraverseTestSpy() { - var goWasCalled = false - var promise: Promise[Int] = Promise[Int]() - val go: () => Promise[Int] = () => { - goWasCalled = true - promise + when(m(Iterable(1, 2, 3))).thenAnswer { + new Answer[Unit] { + def answer(invocation: InvocationOnMock): Unit = { + throw new Exception + } } } - "execute futures in order" in { - val first = new TraverseTestSpy() - val second = new TraverseTestSpy() - val events = Seq(first.go, second.go) + batcher(1) + batcher(2) + batcher(3) // Success here implies no exception was thrown. + } + } - val results = Future.traverseSequentially(events)(f => f()) + batchedTests() + + test( + s"object Future ($name) interruptible should properly ignore the underlying future on interruption") { + val p = Promise[Unit]() + val i = p.interruptible() + val e = new Exception + i.raise(e) + p.setDone() + assert(p.poll.contains(Return(()))) + assert(i.poll.contains(Throw(e))) + } - // At this point, none of the promises - // have been fufilled, so only the first function - // should have been called - assert(first.goWasCalled) - assert(!second.goWasCalled) + test(s"object Future ($name) interruptible should respect the underlying future") { + val p = Promise[Unit]() + val i = p.interruptible() + p.setDone() + assert(p.poll.contains(Return(()))) + assert(i.poll.contains(Return(()))) + } - // once the first promise completes, the next - // function in the sequence should be executed - first.promise.setValue(1) - assert(second.goWasCalled) + test(s"object Future ($name) interruptible should do nothing for const") { + val f = const.value(()) + val i = f.interruptible() + i.raise(new Exception()) + assert(f.poll.contains(Return(()))) + assert(i.poll.contains(Return(()))) + } - // finally, the second promise is fufilled so - // we can Await on and check the results - second.promise.setValue(2) - assert(await(results) == Seq(1, 2)) - } + test(s"object Future ($name) transverseSequentially should execute futures in order") { + val first = new TraverseTestSpy() + val second = new TraverseTestSpy() + val events = Seq(first.go, second.go) - "return with exception when the first future throws" in { - val first = new TraverseTestSpy() - val second = new TraverseTestSpy() - val events = Seq(first.go, second.go) - val results = Future.traverseSequentially(events)(f => f()) + val results = Future.traverseSequentially(events)(f => f()) - first.promise.setException(new Exception) + // At this point, none of the promises + // have been fufilled, so only the first function + // should have been called + assert(first.goWasCalled) + assert(!second.goWasCalled) - intercept[Exception] { await(results) } + // once the first promise completes, the next + // function in the sequence should be executed + first.promise.setValue(1) + assert(second.goWasCalled) - // Since first returned an exception, second should - // never have been called - assert(!second.goWasCalled) - } - } + // finally, the second promise is fufilled so + // we can Await on and check the results + second.promise.setValue(2) + assert(await(results) == Seq(1, 2)) + } - "collect" should { - trait CollectHelper { - val p0, p1 = new HandledPromise[Int] - val f: Future[Seq[Int]] = Future.collect(Seq(p0, p1)) - assert(!f.isDefined) - } + test( + s"object Future ($name) transverseSequentially should return with exception when the first future throws") { + val first = new TraverseTestSpy() + val second = new TraverseTestSpy() + val events = Seq(first.go, second.go) + val results = Future.traverseSequentially(events)(f => f()) - "only return when both futures complete" in { - new CollectHelper { - p0() = Return(1) - assert(!f.isDefined) - p1() = Return(2) - assert(f.isDefined) - assert(await(f) == Seq(1, 2)) - } - } + first.promise.setException(new Exception) - "return with exception if the first future throws" in { - new CollectHelper { - p0() = Throw(new Exception) - intercept[Exception] { await(f) } - } - } + intercept[Exception] { await(results) } - "return with exception if the second future throws" in { - new CollectHelper { - p0() = Return(1) - assert(!f.isDefined) - p1() = Throw(new Exception) - intercept[Exception] { await(f) } - } - } + // Since first returned an exception, second should + // never have been called + assert(!second.goWasCalled) + } - "propagate interrupts" in { - new CollectHelper { - val ps = Seq(p0, p1) - assert(ps.count(_.handled.isDefined) == 0) - f.raise(new Exception) - assert(ps.count(_.handled.isDefined) == 2) - } - } + test(s"object Future ($name) collect should only return when both futures complete") { + new CollectHelper { + p0() = Return(1) + assert(!f.isDefined) + p1() = Return(2) + assert(f.isDefined) + assert(await(f) == Seq(1, 2)) + } + } - "accept maps of futures" in { - val map = Map( - "1" -> Future.value("1"), - "2" -> Future.value("2") - ) + test(s"object Future ($name) collect should return with exception if the first future throws") { + new CollectHelper { + p0() = Throw(new Exception) + intercept[Exception] { await(f) } + } + } - assert(await(Future.collect(map)) == Map("1" -> "1", "2" -> "2")) - } + test( + s"object Future ($name) collect should return with exception if the second future throws") { + new CollectHelper { + p0() = Return(1) + assert(!f.isDefined) + p1() = Throw(new Exception) + intercept[Exception] { await(f) } + } + } - "work correctly if the given map is empty" in { - val map = Map.empty[String, Future[String]] - assert(await(Future.collect(map)).isEmpty) - } + test(s"object Future ($name) collect should propagate interrupts") { + new CollectHelper { + val ps = Seq(p0, p1) + assert(ps.count(_.handled.isDefined) == 0) + f.raise(new Exception) + assert(ps.count(_.handled.isDefined) == 2) + } + } - "return future exception if one of the map values is future exception" in { - val map = Map( - "1" -> Future.value("1"), - "2" -> Future.exception(new Exception) - ) + test(s"object Future ($name) collect should accept maps of futures") { + val map = Map( + "1" -> Future.value("1"), + "2" -> Future.value("2") + ) - intercept[Exception] { - await(Future.collect(map)) - } - } - } + assert(await(Future.collect(map)) == Map("1" -> "1", "2" -> "2")) + } - "collectToTry" should { + test(s"object Future ($name) collect should work correctly if the given map is empty") { + val map = Map.empty[String, Future[String]] + assert(await(Future.collect(map)).isEmpty) + } - trait CollectToTryHelper { - val p0, p1 = new HandledPromise[Int] - val f: Future[Seq[Try[Int]]] = Future.collectToTry(Seq(p0, p1)) - assert(!f.isDefined) - } + test( + s"object Future ($name) collect should return future exception if one of the map values is future exception") { + val map = Map( + "1" -> Future.value("1"), + "2" -> Future.exception(new Exception) + ) - "only return when both futures complete" in { - new CollectToTryHelper { - p0() = Return(1) - assert(!f.isDefined) - p1() = Return(2) - assert(f.isDefined) - assert(await(f) == Seq(Return(1), Return(2))) - } - } + intercept[Exception] { + await(Future.collect(map)) + } + } - "be undefined if the first future throws and the second is undefined" in { - new CollectToTryHelper { - p0() = Throw(new Exception) - assert(!f.isDefined) - } - } + test(s"object Future ($name) collectToTry should only return when both futures complete") { + new CollectToTryHelper { + p0() = Return(1) + assert(!f.isDefined) + p1() = Return(2) + assert(f.isDefined) + assert(await(f) == Seq(Return(1), Return(2))) + } + } - "return both results if the first is defined second future throws" in { - new CollectToTryHelper { - val ex = new Exception - p0() = Return(1) - assert(!f.isDefined) - p1() = Throw(ex) - assert(await(f) == Seq(Return(1), Throw(ex))) - } - } + test( + s"object Future ($name) collectToTry should be undefined if the first future throws and the second is undefined") { + new CollectToTryHelper { + p0() = Throw(new Exception) + assert(!f.isDefined) + } + } - "propagate interrupts" in { - new CollectToTryHelper { - val ps = Seq(p0, p1) - assert(ps.count(_.handled.isDefined) == 0) - f.raise(new Exception) - assert(ps.count(_.handled.isDefined) == 2) - } - } + test( + s"object Future ($name) collectToTry should return both results if the first is defined second future throws") { + new CollectToTryHelper { + val ex = new Exception + p0() = Return(1) + assert(!f.isDefined) + p1() = Throw(ex) + assert(await(f) == Seq(Return(1), Throw(ex))) + } + } + + test(s"object Future ($name) collectToTry should propagate interrupts") { + new CollectToTryHelper { + val ps = Seq(p0, p1) + assert(ps.count(_.handled.isDefined) == 0) + f.raise(new Exception) + assert(ps.count(_.handled.isDefined) == 2) } + } - "propagate locals, restoring original context" in { - val local = new Local[Int] - val f = const.value(111) + test(s"object Future ($name) should propagate locals, restoring original context") { + val local = new Local[Int] + val f = const.value(111) - var ran = 0 - local() = 1010 + var ran = 0 + local() = 1010 + f.ensure { + assert(local().contains(1010)) + local() = 1212 f.ensure { - assert(local().contains(1010)) - local() = 1212 - f.ensure { - assert(local().contains(1212)) - local() = 1313 - ran += 1 - } assert(local().contains(1212)) + local() = 1313 ran += 1 } - - assert(local().contains(1010)) - assert(ran == 2) + assert(local().contains(1212)) + ran += 1 } - "delay execution" in { - val f = const.value(111) + assert(local().contains(1010)) + assert(ran == 2) + } - var count = 0 - f.onSuccess { _ => - assert(count == 0) - f.ensure { - assert(count == 1) - count += 1 - } + test(s"object Future ($name) should delay execution") { + val f = const.value(111) - assert(count == 0) + var count = 0 + f.onSuccess { _ => + assert(count == 0) + f.ensure { + assert(count == 1) count += 1 } - assert(count == 2) + assert(count == 0) + count += 1 } - "are monitored" in { - val inner = const.value(123) - val exc = new Exception("a raw exception") + assert(count == 2) + } - val f = Future.monitored { - inner.ensure { throw exc } - } + test(s"object Future ($name) should are monitored") { + val inner = const.value(123) + val exc = new Exception("a raw exception") - assert(f.poll.contains(Throw(exc))) + val f = Future.monitored { + inner.ensure { throw exc } } - } - s"Future ($name)" should { - "select" which { - trait SelectHelper { - var nhandled = 0 - val p0, p1 = new HandledPromise[Int] - val f: Future[Int] = p0.select(p1) - assert(!f.isDefined) - } + assert(f.poll.contains(Throw(exc))) + } - "select the first [result] to complete" in { - new SelectHelper { - p0() = Return(1) - p1() = Return(2) - assert(await(f) == 1) - } - } + test(s"Future ($name) select should select the first [result] to complete") { + new SelectHelper { + p0() = Return(1) + p1() = Return(2) + assert(await(f) == 1) + } + } - "select the first [exception] to complete" in { - new SelectHelper { - p0() = Throw(new Exception) - p1() = Return(2) - intercept[Exception] { await(f) } - } - } + test(s"Future ($name) select should select the first [exception] to complete") { + new SelectHelper { + p0() = Throw(new Exception) + p1() = Return(2) + intercept[Exception] { await(f) } + } + } - "propagate interrupts" in { - new SelectHelper { - val ps = Seq(p0, p1) - assert(!ps.exists(_.handled.isDefined)) - f.raise(new Exception) - assert(ps.forall(_.handled.isDefined)) - } - } + test(s"Future ($name) select should propagate interrupts") { + new SelectHelper { + val ps = Seq(p0, p1) + assert(!ps.exists(_.handled.isDefined)) + f.raise(new Exception) + assert(ps.forall(_.handled.isDefined)) } + } - def testJoin(label: String, joiner: ((Future[Int], Future[Int]) => Future[(Int, Int)])): Unit = { - s"join($label)" should { - trait JoinHelper { - val p0 = new HandledPromise[Int] - val p1 = new HandledPromise[Int] - val f = joiner(p0, p1) - assert(!f.isDefined) - } + def testJoin( + label: String, + joiner: ((Future[Int], Future[Int]) => Future[(Int, Int)]) + ): Unit = { + test(s"Future ($name) join($label) should propagate interrupts") { + trait JoinHelper { + val p0 = new HandledPromise[Int] + val p1 = new HandledPromise[Int] + val f = joiner(p0, p1) + assert(!f.isDefined) + } - for { - left <- Seq(Ret(1), Thr(new Exception), Nvr) - right <- Seq(Ret(2), Thr(new Exception), Nvr) - leftFirst <- Seq(true, false) - } { - new JoinHelper { - if (leftFirst) { - satisfy(p0, left) - satisfy(p1, right) - } else { - satisfy(p1, right) - satisfy(p0, left) - } - (left, right) match { - // Two Throws are special because leftFirst determines the - // exception (e0 or e1). Otherwise, every other case with a - // Throw will result in just one exception. - case (Thr(e0), Thr(e1)) => - val actual = intercept[Exception] { await(f) } - assert(actual == (if (leftFirst) e0 else e1), explain(left, right, leftFirst)) - case (_, Thr(exn)) => - val actual = intercept[Exception] { await(f) } - assert(actual == exn, explain(left, right, leftFirst)) - case (Thr(exn), _) => - val actual = intercept[Exception] { await(f) } - assert(actual == exn, explain(left, right, leftFirst)) - case (Nvr, Ret(_)) | (Ret(_), Nvr) | (Nvr, Nvr) => assert(!f.isDefined) - case (Ret(a), Ret(b)) => assert(await(f) == (a, b), explain(left, right, leftFirst)) - } + for { + left <- Seq(Ret(1), Thr(new Exception), Nvr) + right <- Seq(Ret(2), Thr(new Exception), Nvr) + leftFirst <- Seq(true, false) + } { + new JoinHelper { + if (leftFirst) { + satisfy(p0, left) + satisfy(p1, right) + } else { + satisfy(p1, right) + satisfy(p0, left) } - } - - "propagate interrupts" in { - new JoinHelper { - assert(p0.handled.isEmpty) - assert(p1.handled.isEmpty) - val exc = new Exception - f.raise(exc) - assert(p0.handled.contains(exc)) - assert(p1.handled.contains(exc)) + (left, right) match { + // Two Throws are special because leftFirst determines the + // exception (e0 or e1). Otherwise, every other case with a + // Throw will result in just one exception. + case (Thr(e0), Thr(e1)) => + val actual = intercept[Exception] { await(f) } + assert(actual == (if (leftFirst) e0 else e1), explain(left, right, leftFirst)) + case (_, Thr(exn)) => + val actual = intercept[Exception] { await(f) } + assert(actual == exn, explain(left, right, leftFirst)) + case (Thr(exn), _) => + val actual = intercept[Exception] { await(f) } + assert(actual == exn, explain(left, right, leftFirst)) + case (Nvr, Ret(_)) | (Ret(_), Nvr) | (Nvr, Nvr) => assert(!f.isDefined) + case (Ret(a), Ret(b)) => + val expected: (Int, Int) = (a, b) + val result = await(f) + val isEqual = result == expected + assert(isEqual, explain(left, right, leftFirst)) } } } + + new JoinHelper { + assert(p0.handled.isEmpty) + assert(p1.handled.isEmpty) + val exc = new Exception + f.raise(exc) + assert(p0.handled.contains(exc)) + assert(p1.handled.contains(exc)) + } } + } - testJoin("f join g", _ join _) - testJoin("Future.join(f, g)", Future.join(_, _)) + testJoin("f join g", _ join _) + testJoin("Future.join(f, g)", Future.join(_, _)) - "toJavaFuture" should { - "return the same thing as our Future when initialized" which { - val f = const.value(1) - val jf = f.toJavaFuture - assert(await(f) == jf.get()) - "must both be done" in { - assert(f.isDefined) - assert(jf.isDone) - assert(!jf.isCancelled) - } - } + def testJavaFuture(methodName: String, fn: Future[Int] => JFuture[_ <: Int]): Unit = { + test( + s"Future ($name) $methodName should return the same thing as our Future when initialized") { + val f = const.value(1) + val jf = fn(f) + assert(await(f) == jf.get()) + assert(f.isDefined) + assert(jf.isDone) + assert(!jf.isCancelled) + } - "return the same thing as our Future when set later" which { - val f = new Promise[Int] - val jf = f.toJavaFuture - f.setValue(1) - assert(await(f) == jf.get()) - "must both be done" in { - assert(f.isDefined) - assert(jf.isDone) - assert(!jf.isCancelled) - } - } + test( + s"Future ($name) $methodName should return the same thing as our Future when set later") { + val f = new Promise[Int] + val jf = fn(f) + f.setValue(1) + assert(await(f) == jf.get()) + assert(f.isDefined) + assert(jf.isDone) + assert(!jf.isCancelled) + } - "java future should throw an exception" in { - val f = new Promise[Int] - val jf = f.toJavaFuture - val e = new RuntimeException() - f.setException(e) - val actual = intercept[RuntimeException] { jf.get() } - assert(actual == e) - } + test(s"Future ($name) $methodName java future should throw an exception when failed later") { + val f = new Promise[Int] + val jf = fn(f) + val e = new RuntimeException() + f.setException(e) + val actual = intercept[ExecutionException] { jf.get() } + val cause = intercept[RuntimeException] { throw actual.getCause } + assert(cause == e) + } - "interrupt Future when cancelled" in { - val f = new HandledPromise[Int] - val jf = f.toJavaFuture - assert(f.handled.isEmpty) - jf.cancel(true) - assert(f.handled match { - case Some(_: java.util.concurrent.CancellationException) => true - case _ => false - }) - } + test(s"Future ($name) $methodName java future should throw an exception when failed early") { + val f = new Promise[Int] + val e = new RuntimeException() + f.setException(e) + val jf = fn(f) + val actual = intercept[ExecutionException] { jf.get() } + val cause = intercept[RuntimeException] { throw actual.getCause } + assert(cause == e) } - "monitored" should { - trait MonitoredHelper { - val inner = new HandledPromise[Int] - val exc = new Exception("some exception") + test(s"Future ($name) $methodName should interrupt Future when cancelled") { + val f = new HandledPromise[Int] + val jf = fn(f) + assert(f.handled.isEmpty) + jf.cancel(true) + intercept[java.util.concurrent.CancellationException] { + throw f.handled.get } + } + } - "catch raw exceptions (direct)" in { - new MonitoredHelper { - val f = Future.monitored { - throw exc - inner - } - assert(f.poll.contains(Throw(exc))) - } + testJavaFuture("toJavaFuture", { (f: Future[Int]) => f.toJavaFuture }) + testJavaFuture("toCompletableFuture", { (f: Future[Int]) => f.toCompletableFuture }) + + test(s"Future ($name) monitored should catch raw exceptions (direct)") { + new MonitoredHelper { + val f = Future.monitored { + throw exc + inner } + assert(f.poll.contains(Throw(exc))) + } + } - "catch raw exceptions (indirect), interrupting computation" in { - new MonitoredHelper { - val inner1 = new Promise[Int] - var ran = false - val f: Future[Int] = Future.monitored { - inner1.ensure { - // Note that these are sequenced so that interrupts - // will be delivered before inner's handler is cleared. - ran = true - try { - inner.update(Return(1)) - } catch { - case _: Throwable => fail() - } - }.ensure { - throw exc + test( + s"Future ($name) monitored should catch raw exceptions (indirect), interrupting computation") { + new MonitoredHelper { + val inner1 = new Promise[Int] + var ran = false + val f: Future[Int] = Future.monitored { + inner1 + .ensure { + // Note that these are sequenced so that interrupts + // will be delivered before inner's handler is cleared. + ran = true + try { + inner.update(Return(1)) + } catch { + case _: Throwable => fail() } - inner } - assert(!ran) - assert(f.poll.isEmpty) - assert(inner.handled.isEmpty) - inner1.update(Return(1)) - assert(ran) - assert(inner.isDefined) - assert(f.poll.contains(Throw(exc))) - - assert(inner.handled.contains(exc)) - } - } - - "link" in { - new MonitoredHelper { - val f: Future[Int] = Future.monitored { inner } - assert(inner.handled.isEmpty) - f.raise(exc) - assert(inner.handled.contains(exc)) - } - } + .ensure { + throw exc + } + inner + } + assert(!ran) + assert(f.poll.isEmpty) + assert(inner.handled.isEmpty) + inner1.update(Return(1)) + assert(ran) + assert(inner.isDefined) + assert(f.poll.contains(Throw(exc))) - "doesn't leak the underlying promise after completion" in { - new MonitoredHelper { - val inner1 = new Promise[String] - val inner2 = new Promise[String] - val f: Future[String] = Future.monitored { inner2.ensure(()); inner1 } - val s: String = "." * 1024 - val sSize: Long = ObjectSizeCalculator.getObjectSize(s) - inner1.setValue(s) - val inner2Size: Long = ObjectSizeCalculator.getObjectSize(inner2) - assert(inner2Size < sSize) - } - } + assert(inner.handled.contains(exc)) } } - s"Promise ($name)" should { - "apply" which { - "when we're inside of a respond block (without deadlocking)" in { - val f = Future(1) - var didRun = false - f.foreach { _ => - f mustProduce Return(1) - didRun = true - } - assert(didRun) - } + test(s"Future ($name) monitored should link") { + new MonitoredHelper { + val f: Future[Int] = Future.monitored { inner } + assert(inner.handled.isEmpty) + f.raise(exc) + assert(inner.handled.contains(exc)) } + } - "map" which { - "when it's all chill" in { - val f = Future(1).map { x => - x + 1 - } - assert(await(f) == 2) - } - - "when there's a problem in the passed in function" in { - val e = new Exception - val f = Future(1).map { x => - throw e - x + 1 - } - val actual = intercept[Exception] { - await(f) - } - assert(actual == e) + // we only know this works as expected when running with JDK 8 + if (System.getProperty("java.version").startsWith("1.8")) + test(s"Future ($name) monitored should not leak the underlying promise after completion") { + new MonitoredHelper { + val inner1 = new Promise[String] + val inner2 = new Promise[String] + val f: Future[String] = Future.monitored { inner2.ensure(()); inner1 } + val s: String = "." * 1024 * 10 + val sSize: Long = ObjectSizeCalculator.getObjectSize(s) + inner1.setValue(s) + val inner2Size: Long = ObjectSizeCalculator.getObjectSize(inner2) + assert(inner2Size < sSize) } } - "transform" should { - val e = new Exception("rdrr") + test( + s"Promise ($name) should apply when we're inside of a respond block (without deadlocking)") { + val f = Future(1) + var didRun = false + f.foreach { _ => + f mustProduce Return(1) + didRun = true + } + assert(didRun) + } - "values" in { - const.value(1).transform { - case Return(v) => const.value(v + 1) - case Throw(_) => const.value(0) - }.mustProduce(Return(2)) - } + test(s"Promise ($name) should map when it's all chill") { + val f = Future(1).map { x => x + 1 } + assert(await(f) == 2) + } - "exceptions" in { - const.exception(e).transform { - case Return(_) => const.value(1) - case Throw(_) => const.value(0) - }.mustProduce(Return(0)) - } + test(s"Promise ($name) should map when there's a problem in the passed in function") { + val e = new Exception + val f = Future(1).map { x => + throw e + x + 1 + } + val actual = intercept[Exception] { + await(f) + } + assert(actual == e) + } - "exceptions thrown during transformation" in { - const.value(1).transform { - case Return(_) => const.value(throw e) - case Throw(_) => const.value(0) - }.mustProduce(Throw(e)) + test(s"Promise ($name) should transform values") { + val e = new Exception("rdrr") + const + .value(1) + .transform { + case Return(v) => const.value(v + 1) + case Throw(_) => const.value(0) } + .mustProduce(Return(2)) + } - "non local returns executed during transformation" in { - def ret(): String = { - val f = const.value(1).transform { - case Return(_) => - val fn = { () => - return "OK" - } - fn() - Future.value(ret()) - case Throw(_) => const.value(0) - } - assert(f.poll.isDefined) - val e = intercept[FutureNonLocalReturnControl] { - f.poll.get.get() - } + test(s"Promise ($name) should transform exceptions") { + val e = new Exception("rdrr") - val g = e.getCause match { - case t: NonLocalReturnControl[_] => t.asInstanceOf[NonLocalReturnControl[String]] - case _ => - fail() - } - assert(g.value == "OK") - "bleh" - } - ret() + const + .exception(e) + .transform { + case Return(_) => const.value(1) + case Throw(_) => const.value(0) } + .mustProduce(Return(0)) + } - "fatal exceptions thrown during transformation" in { - val e = new FatalException() + test(s"Promise ($name) should transform exceptions thrown during transformation") { + val e = new Exception("rdrr") - val actual = intercept[FatalException] { - const.value(1).transform { - case Return(_) => const.value(throw e) - case Throw(_) => const.value(0) - } - } - assert(actual == e) + const + .value(1) + .transform { + case Return(_) => const.value(throw e) + case Throw(_) => const.value(0) } + .mustProduce(Throw(e)) + } - "monitors fatal exceptions" in { - val m = new HandledMonitor() - val exc = new FatalException() - - assert(m.handled == null) + test(s"Promise ($name) should transform non local returns executed during transformation") { + val e = new Exception("rdrr") - val actual = intercept[FatalException] { - Monitor.using(m) { - const.value(1).transform { _ => throw exc } - } - } + def ret(): String = { + val f = const.value(1).transform { + case Return(_) => + val fn = { () => return "OK" } + fn() + Future.value(ret()) + case Throw(_) => const.value(0) + } + assert(f.poll.isDefined) + val e = intercept[FutureNonLocalReturnControl] { + f.poll.get.get() + } - assert(actual == exc) - assert(m.handled == exc) + val g = e.getCause match { + case t: NonLocalReturnControl[_] => t.asInstanceOf[NonLocalReturnControl[String]] + case _ => + fail() } + assert(g.value == "OK") + "bleh" } + ret() + } - "transformedBy" should { - val e = new Exception("rdrr") + test(s"Promise ($name) should transform fatal exceptions thrown during transformation") { + val e = new FatalException() - "flatMap" in { - const - .value(1) - .transformedBy(new FutureTransformer[Int, Int] { - override def flatMap(value: Int): Future[Int] = const.value(value + 1) - override def rescue(t: Throwable): Future[Int] = const.value(0) - }).mustProduce(Return(2)) + val actual = intercept[FatalException] { + const.value(1).transform { + case Return(_) => const.value(throw e) + case Throw(_) => const.value(0) } + } + assert(actual == e) + } - "rescue" in { - const - .exception(e) - .transformedBy(new FutureTransformer[Int, Int] { - override def flatMap(value: Int): Future[Int] = const.value(value + 1) - override def rescue(t: Throwable): Future[Int] = const.value(0) - }).mustProduce(Return(0)) - } + test(s"Promise ($name) transform monitors fatal exceptions") { + val e = new Exception("rdrr") - "exceptions thrown during transformation" in { - const - .value(1) - .transformedBy(new FutureTransformer[Int, Int] { - override def flatMap(value: Int): Future[Int] = throw e - override def rescue(t: Throwable): Future[Int] = const.value(0) - }).mustProduce(Throw(e)) - } + val m = new HandledMonitor() + val exc = new FatalException() - "map" in { - const - .value(1) - .transformedBy(new FutureTransformer[Int, Int] { - override def map(value: Int): Int = value + 1 - override def handle(t: Throwable): Int = 0 - }).mustProduce(Return(2)) - } + assert(m.handled == null) - "handle" in { - const - .exception(e) - .transformedBy(new FutureTransformer[Int, Int] { - override def map(value: Int): Int = value + 1 - override def handle(t: Throwable): Int = 0 - }).mustProduce(Return(0)) + val actual = intercept[FatalException] { + Monitor.using(m) { + const.value(1).transform { _ => throw exc } } } - def testSequence(which: String, seqop: (Future[Unit], () => Future[Unit]) => Future[Unit]): Unit = { - which when { - "successes" should { - "interruption of the produced future" which { - "before the antecedent Future completes, propagates back to the antecedent" in { - val f1, f2 = new HandledPromise[Unit] - val f3 = seqop(f1, () => f2) - assert(f1.handled.isEmpty) - assert(f2.handled.isEmpty) - f3.raise(new Exception) - assert(f1.handled.isDefined) - assert(f2.handled.isEmpty) - f1() = Return.Unit - assert(f2.handled.isDefined) - } + assert(actual == exc) + assert(m.handled == exc) + } - "after the antecedent Future completes, does not propagate back to the antecedent" in { - val f1, f2 = new HandledPromise[Unit] - val f3 = seqop(f1, () => f2) - assert(f1.handled.isEmpty) - assert(f2.handled.isEmpty) - f1() = Return.Unit - f3.raise(new Exception) - assert(f1.handled.isEmpty) - assert(f2.handled.isDefined) - } + test(s"Promise ($name) transformedBy flatMap") { + const + .value(1) + .transformedBy(new FutureTransformer[Int, Int] { + override def flatMap(value: Int): Future[Int] = const.value(value + 1) + override def rescue(t: Throwable): Future[Int] = const.value(0) + }) + .mustProduce(Return(2)) + } - "forward through chains" in { - val f1, f2 = new Promise[Unit] - val exc = new Exception - val f3 = new Promise[Unit] - var didInterrupt = false - f3.setInterruptHandler { - case `exc` => didInterrupt = true - } - val f4 = seqop(f1, () => seqop(f2, () => f3)) - f4.raise(exc) - assert(!didInterrupt) - f1.setDone() - assert(!didInterrupt) - f2.setDone() - assert(didInterrupt) - } - } - } + test(s"Promise ($name) transformedBy rescue") { + val e = new Exception("rdrr") + const + .exception(e) + .transformedBy(new FutureTransformer[Int, Int] { + override def flatMap(value: Int): Future[Int] = const.value(value + 1) + override def rescue(t: Throwable): Future[Int] = const.value(0) + }) + .mustProduce(Return(0)) + } - "failures" should { - val e = new Exception - val g = seqop(Future[Unit](throw e), () => Future.Done) + test(s"Promise ($name) transformedBy exceptions thrown during transformation") { + val e = new Exception("rdrr") + const + .value(1) + .transformedBy(new FutureTransformer[Int, Int] { + override def flatMap(value: Int): Future[Int] = throw e + override def rescue(t: Throwable): Future[Int] = const.value(0) + }) + .mustProduce(Throw(e)) + } - "apply" in { - val actual = intercept[Exception] { await(g) } - assert(actual == e) - } + test(s"Promise ($name) transformedBy map") { + const + .value(1) + .transformedBy(new FutureTransformer[Int, Int] { + override def map(value: Int): Int = value + 1 + override def handle(t: Throwable): Int = 0 + }) + .mustProduce(Return(2)) + } - "respond" in { - g mustProduce Throw(e) - } + test(s"Promise ($name) transformedBy handle") { + val e = new Exception("rdrr") + const + .exception(e) + .transformedBy(new FutureTransformer[Int, Int] { + override def map(value: Int): Int = value + 1 + override def handle(t: Throwable): Int = 0 + }) + .mustProduce(Return(0)) + } - "when there is an exception in the passed in function" in { - val e = new Exception - val f = seqop(Future.Done, () => throw e) - val actual = intercept[Exception] { await(f) } - assert(actual == e) - } - } - } + def testSequence( + which: String, + seqop: (Future[Unit], () => Future[Unit]) => Future[Unit] + ): Unit = { + test( + s"Promise ($name) $which successes: interruption of the produced future, before the antecedent Future completes, propagates back to the antecedent") { + val f1, f2 = new HandledPromise[Unit] + val f3 = seqop(f1, () => f2) + assert(f1.handled.isEmpty) + assert(f2.handled.isEmpty) + f3.raise(new Exception) + assert(f1.handled.isDefined) + assert(f2.handled.isEmpty) + f1() = Return.Unit + assert(f2.handled.isDefined) } - testSequence( - "flatMap", - (a, next) => - a.flatMap { _ => - next() - } - ) - testSequence("before", (a, next) => a.before { next() }) + test( + s"Promise ($name) $which successes: interruption of the produced future, after the antecedent Future completes, does not propagate back to the antecedent") { + val f1, f2 = new HandledPromise[Unit] + val f3 = seqop(f1, () => f2) + assert(f1.handled.isEmpty) + assert(f2.handled.isEmpty) + f1() = Return.Unit + f3.raise(new Exception) + assert(f1.handled.isEmpty) + assert(f2.handled.isDefined) + } - "flatMap (values)" should { - val f = Future(1).flatMap { x => - Future(x + 1) - } + test( + s"Promise ($name) $which successes: interruption of the produced future, forward through chains") { + val f1, f2 = new Promise[Unit] + val exc = new Exception + val f3 = new Promise[Unit] + var didInterrupt = false + f3.setInterruptHandler { + case `exc` => didInterrupt = true + } + val f4 = seqop(f1, () => seqop(f2, () => f3)) + f4.raise(exc) + assert(!didInterrupt) + f1.setDone() + assert(!didInterrupt) + f2.setDone() + assert(didInterrupt) + } - "apply" in { - assert(await(f) == 2) - } + test(s"Promise ($name) $which failures shouldapply") { + val e = new Exception + val g = seqop(Future[Unit](throw e), () => Future.Done) + val actual = intercept[Exception] { await(g) } + assert(actual == e) + } - "respond" in { - f mustProduce Return(2) - } + test(s"Promise ($name) $which failures shouldrespond") { + val e = new Exception + val g = seqop(Future[Unit](throw e), () => Future.Done) + g mustProduce Throw(e) } - "flatten" should { - "successes" in { - val f = Future(Future(1)) - f.flatten mustProduce Return(1) - } + test( + s"Promise ($name) $which failures should when there is an exception in the passed in function") { + val e = new Exception + val f = seqop(Future.Done, () => throw e) + val actual = intercept[Exception] { await(f) } + assert(actual == e) + } + } - "shallow failures" in { - val e = new Exception - val f: Future[Future[Int]] = const.exception(e) - f.flatten mustProduce Throw(e) - } + testSequence( + "flatMap", + (a, next) => a.flatMap { _ => next() } + ) + testSequence("before", (a, next) => a.before { next() }) - "deep failures" in { - val e = new Exception - val f: Future[Future[Int]] = const.value(const.exception(e)) - f.flatten mustProduce Throw(e) - } + test(s"Promise ($name) flatMap (values) apply") { + val f = Future(1).flatMap { x => Future(x + 1) } + assert(await(f) == 2) + } - "interruption" in { - val f1 = new HandledPromise[Future[Int]] - val f2 = new HandledPromise[Int] - val f = f1.flatten - assert(f1.handled.isEmpty) - assert(f2.handled.isEmpty) - f.raise(new Exception) - f1.handled match { - case Some(_) => - case None => fail() - } - assert(f2.handled.isEmpty) - f1() = Return(f2) - f2.handled match { - case Some(_) => - case None => fail() - } - } - } + test(s"Promise ($name) flatMap (values) respond") { + val f = Future(1).flatMap { x => Future(x + 1) } + f mustProduce Return(2) + } - "rescue" should { - val e = new Exception + test(s"Promise ($name) flatten successes") { + val f = Future(Future(1)) + f.flatten mustProduce Return(1) + } - "successes" which { - val f = Future(1).rescue { case _ => Future(2) } + test(s"Promise ($name) flatten shallow failures") { + val e = new Exception + val f: Future[Future[Int]] = const.exception(e) + f.flatten mustProduce Throw(e) + } - "apply" in { - assert(await(f) == 1) - } + test(s"Promise ($name) flatten deep failures") { + val e = new Exception + val f: Future[Future[Int]] = const.value(const.exception(e)) + f.flatten mustProduce Throw(e) + } - "respond" in { - f mustProduce Return(1) - } - } + test(s"Promise ($name) flatten interruption") { + val f1 = new HandledPromise[Future[Int]] + val f2 = new HandledPromise[Int] + val f = f1.flatten + assert(f1.handled.isEmpty) + assert(f2.handled.isEmpty) + f.raise(new Exception) + f1.handled match { + case Some(_) => + case None => fail() + } + assert(f2.handled.isEmpty) + f1() = Return(f2) + f2.handled match { + case Some(_) => + case None => fail() + } + } - "failures" which { - val g = Future[Int](throw e).rescue { case _ => Future(2) } + test(s"Promise ($name) rescue successses apply") { + val f = Future(1).rescue { case _ => Future(2) } + assert(await(f) == 1) + } - "apply" in { - assert(await(g) == 2) - } + test(s"Promise ($name) rescue successes: respond") { + val f = Future(1).rescue { case _ => Future(2) } + f mustProduce Return(1) + } - "respond" in { - g mustProduce Return(2) - } + test(s"Promise ($name) rescue failures: apply") { + val e = new Exception + val g = Future[Int](throw e).rescue { case _ => Future(2) } + assert(await(g) == 2) + } - "when the error handler errors" in { - val g = Future[Int](throw e).rescue { case x => throw x; Future(2) } - val actual = intercept[Exception] { await(g) } - assert(actual == e) - } - } + test(s"Promise ($name) rescue failures: respond") { + val e = new Exception + val g = Future[Int](throw e).rescue { case _ => Future(2) } + g mustProduce Return(2) + } - "interruption of the produced future" which { - "before the antecedent Future completes, propagates back to the antecedent" in { - val f1, f2 = new HandledPromise[Int] - val f = f1.rescue { case _ => f2 } - assert(f1.handled.isEmpty) - assert(f2.handled.isEmpty) - f.raise(new Exception) - f1.handled match { - case Some(_) => - case None => fail() - } - assert(f2.handled.isEmpty) - f1() = Throw(new Exception) - f2.handled match { - case Some(_) => - case None => fail() - } - } + test(s"Promise ($name) rescue failures: when the error handler errors") { + val e = new Exception + val g = Future[Int](throw e).rescue { case x => throw x; Future(2) } + val actual = intercept[Exception] { await(g) } + assert(actual == e) + } - "after the antecedent Future completes, does not propagate back to the antecedent" in { - val f1, f2 = new HandledPromise[Int] - val f = f1.rescue { case _ => f2 } - assert(f1.handled.isEmpty) - assert(f2.handled.isEmpty) - f1() = Throw(new Exception) - f.raise(new Exception) - assert(f1.handled.isEmpty) - f2.handled match { - case Some(_) => - case None => fail() - } - } - } + test( + s"Promise ($name) rescue interruption of the produced future before the antecedent Future completes, propagates back to the antecedent") { + val f1, f2 = new HandledPromise[Int] + val f = f1.rescue { case _ => f2 } + assert(f1.handled.isEmpty) + assert(f2.handled.isEmpty) + f.raise(new Exception) + f1.handled match { + case Some(_) => + case None => fail() } - - "foreach" in { - var wasCalledWith: Option[Int] = None - val f = Future(1) - f.foreach { i => - wasCalledWith = Some(i) - } - assert(wasCalledWith.contains(1)) + assert(f2.handled.isEmpty) + f1() = Throw(new Exception) + f2.handled match { + case Some(_) => + case None => fail() } + } - "respond" should { - "when the result has arrived" in { - var wasCalledWith: Option[Int] = None - val f = Future(1) - f.respond { - case Return(i) => wasCalledWith = Some(i) - case Throw(e) => fail(e.toString) - } - assert(wasCalledWith.contains(1)) - } + test( + s"Promise ($name) rescue interruption of the produced future after the antecedent Future completes, does not propagate back to the antecedent") { + val f1, f2 = new HandledPromise[Int] + val f = f1.rescue { case _ => f2 } + assert(f1.handled.isEmpty) + assert(f2.handled.isEmpty) + f1() = Throw(new Exception) + f.raise(new Exception) + assert(f1.handled.isEmpty) + f2.handled match { + case Some(_) => + case None => fail() + } + } - "when the result has not yet arrived it buffers computations" in { - var wasCalledWith: Option[Int] = None - val f = new Promise[Int] - f.foreach { i => - wasCalledWith = Some(i) - } - assert(wasCalledWith.isEmpty) - f() = Return(1) - assert(wasCalledWith.contains(1)) - } + test(s"Promise ($name) foreach") { + var wasCalledWith: Option[Int] = None + val f = Future(1) + f.foreach { i => wasCalledWith = Some(i) } + assert(wasCalledWith.contains(1)) + } - "runs callbacks just once and in order (lifo)" in { - var i, j, k, h = 0 - val p = new Promise[Int] + test(s"Promise ($name) respond when the result has arrived") { + var wasCalledWith: Option[Int] = None + val f = Future(1) + f.respond { + case Return(i) => wasCalledWith = Some(i) + case Throw(e) => fail(e.toString) + } + assert(wasCalledWith.contains(1)) + } - p.ensure { - i = i + j + k + h + 1 - }.ensure { - j = i + j + k + h + 1 - }.ensure { - k = i + j + k + h + 1 - }.ensure { - h = i + j + k + h + 1 - } + test(s"Promise ($name) respond when the result has not yet arrived it buffers computations") { + var wasCalledWith: Option[Int] = None + val f = new Promise[Int] + f.foreach { i => wasCalledWith = Some(i) } + assert(wasCalledWith.isEmpty) + f() = Return(1) + assert(wasCalledWith.contains(1)) + } - assert(i == 0) - assert(j == 0) - assert(k == 0) - assert(h == 0) + test(s"Promise ($name) respond runs callbacks just once and in order (lifo)") { + var i, j, k, h = 0 + val p = new Promise[Int] - p.setValue(1) - assert(i == 8) - assert(j == 4) - assert(k == 2) - assert(h == 1) + p.ensure { + i = i + j + k + h + 1 + } + .ensure { + j = i + j + k + h + 1 + } + .ensure { + k = i + j + k + h + 1 + } + .ensure { + h = i + j + k + h + 1 } - "monitor exceptions" in { - val m = new HandledMonitor() - val exc = new Exception + assert(i == 0) + assert(j == 0) + assert(k == 0) + assert(h == 0) - assert(m.handled == null) + p.setValue(1) + assert(i == 8) + assert(j == 4) + assert(k == 2) + assert(h == 1) + } - Monitor.using(m) { - const.value(1).ensure { throw exc } - } + test(s"Promise ($name) respond monitor exceptions") { + val m = new HandledMonitor() + val exc = new Exception - assert(m.handled == exc) - } - } + assert(m.handled == null) - "willEqual" in { - assert(await(const.value(1).willEqual(const.value(1)), 1.second)) + Monitor.using(m) { + const.value(1).ensure { throw exc } } - "Future() handles exceptions" in { - val e = new Exception - val f = Future[Int] { throw e } - val actual = intercept[Exception] { await(f) } - assert(actual == e) - } + assert(m.handled == exc) + } - "propagate locals" in { - val local = new Local[Int] - val promise0 = new Promise[Unit] - val promise1 = new Promise[Unit] + test(s"Promise ($name) willEqual") { + assert(await(const.value(1).willEqual(const.value(1)), 1.second)) + } - local() = 1010 + test(s"Promise ($name) Future() handles exceptions") { + val e = new Exception + val f = Future[Int] { throw e } + val actual = intercept[Exception] { await(f) } + assert(actual == e) + } - val both = promise0.flatMap { _ => - val local0 = local() - promise1.map { _ => - val local1 = local() - (local0, local1) - } - } + test(s"Promise ($name) propagate locals") { + val local = new Local[Int] + val promise0 = new Promise[Unit] + val promise1 = new Promise[Unit] - local() = 123 - promise0() = Return.Unit - local() = 321 - promise1() = Return.Unit + local() = 1010 - assert(both.isDefined) - assert(await(both) == ((Some(1010), Some(1010)))) + val both = promise0.flatMap { _ => + val local0 = local() + promise1.map { _ => + val local1 = local() + (local0, local1) + } } - "propagate locals across threads" in { - val local = new Local[Int] - val promise = new Promise[Option[Int]] - - local() = 123 - val done = promise.map { otherValue => - (otherValue, local()) - } + local() = 123 + promise0() = Return.Unit + local() = 321 + promise1() = Return.Unit - val t = new Thread { - override def run(): Unit = { - local() = 1010 - promise() = Return(local()) - } - } + assert(both.isDefined) + assert(await(both) == ((Some(1010), Some(1010)))) + } - t.run() - t.join() + test(s"Promise ($name) propagate locals across threads") { + val local = new Local[Int] + val promise = new Promise[Option[Int]] - assert(done.isDefined) - assert(await(done) == ((Some(1010), Some(123)))) - } + local() = 123 + val done = promise.map { otherValue => (otherValue, local()) } - "poll" should { - trait PollHelper { - val p = new Promise[Int] - } - "when waiting" in { - new PollHelper { - assert(p.poll.isEmpty) - } + val t = new Thread { + override def run(): Unit = { + local() = 1010 + promise() = Return(local()) } + } - "when succeeding" in { - new PollHelper { - p.setValue(1) - assert(p.poll.contains(Return(1))) - } - } + t.run() + t.join() - "when failing" in { - new PollHelper { - val e = new Exception - p.setException(e) - assert(p.poll.contains(Throw(e))) - } - } - } + assert(done.isDefined) + assert(await(done) == ((Some(1010), Some(123)))) + } - val uw = (p: Promise[Int], d: Duration, t: Timer) => { - p.within(d)(t) + test(s"Promise ($name) poll when waiting") { + new PollHelper { + assert(p.poll.isEmpty) } + } - val ub = (p: Promise[Int], d: Duration, t: Timer) => { - p.by(d.fromNow)(t) + test(s"Promise ($name) poll when succeeding") { + new PollHelper { + p.setValue(1) + assert(p.poll.contains(Return(1))) } + } - Seq(("within", uw), ("by", ub)).foreach { - case (label, use) => - label should { - "when we run out of time" in { - implicit val timer: Timer = new JavaTimer - val p = new HandledPromise[Int] - intercept[TimeoutException] { await(use(p, 50.milliseconds, timer)) } - timer.stop() - assert(p.handled.isEmpty) - } - - "when everything is chill" in { - implicit val timer: Timer = new JavaTimer - val p = new Promise[Int] - p.setValue(1) - assert(await(use(p, 50.milliseconds, timer)) == 1) - timer.stop() - } - - "when timeout is forever" in { - // We manage to throw an exception inside - // the scala compiler if we use MockTimer - // here. Sigh. - implicit val timer: Timer = FailingTimer - val p = new Promise[Int] - assert(use(p, Duration.Top, timer) == p) - } - - "when future already satisfied" in { - implicit val timer: Timer = new NullTimer - val p = new Promise[Int] - p.setValue(3) - assert(use(p, 1.minute, timer) == p) - } - - "interruption" in Time.withCurrentTimeFrozen { _ => - implicit val timer: Timer = new MockTimer - val p = new HandledPromise[Int] - val f = use(p, 50.milliseconds, timer) - assert(p.handled.isEmpty) - f.raise(new Exception) - p.handled match { - case Some(_) => - case None => fail() - } - } - } + test(s"Promise ($name) poll when failing") { + new PollHelper { + val e = new Exception + p.setException(e) + assert(p.poll.contains(Throw(e))) } + } - "raiseWithin" should { - "when we run out of time" in { - implicit val timer: Timer = new JavaTimer - val p = new HandledPromise[Int] - intercept[TimeoutException] { - await(p.raiseWithin(50.milliseconds)) - } - timer.stop() - p.handled match { - case Some(_) => - case None => fail() - } - } + val uw = (p: Promise[Int], d: Duration, t: Timer) => { + p.within(d)(t) + } - "when we run out of time, throw our stuff" in { - implicit val timer: Timer = new JavaTimer - class SkyFallException extends Exception("let the skyfall") - val skyFall = new SkyFallException - val p = new HandledPromise[Int] - intercept[SkyFallException] { - await(p.raiseWithin(50.milliseconds, skyFall)) - } - timer.stop() - p.handled match { - case Some(_) => - case None => fail() - } - assert(p.handled.contains(skyFall)) - } + val ub = (p: Promise[Int], d: Duration, t: Timer) => { + p.by(d.fromNow)(t) + } - "when we are within timeout, but inner throws TimeoutException, we don't raise" in { + Seq(("within", uw), ("by", ub)).foreach { + case (label, use) => + test(s"Promise ($name) $label when we run out of time") { implicit val timer: Timer = new JavaTimer - class SkyFallException extends Exception("let the skyfall") - val skyFall = new SkyFallException val p = new HandledPromise[Int] - intercept[TimeoutException] { - await( - p.within(20.milliseconds).raiseWithin(50.milliseconds, skyFall) - ) - } + intercept[TimeoutException] { await(use(p, 50.milliseconds, timer)) } timer.stop() assert(p.handled.isEmpty) } - "when everything is chill" in { + test(s"Promise ($name) $label when everything is chill") { implicit val timer: Timer = new JavaTimer val p = new Promise[Int] p.setValue(1) - assert(await(p.raiseWithin(50.milliseconds)) == 1) + assert(await(use(p, 50.milliseconds, timer)) == 1) timer.stop() } - "when timeout is forever" in { + test(s"Promise ($name) $label when timeout is forever") { // We manage to throw an exception inside // the scala compiler if we use MockTimer // here. Sigh. implicit val timer: Timer = FailingTimer val p = new Promise[Int] - assert(p.raiseWithin(Duration.Top) == p) + assert(use(p, Duration.Top, timer) == p) } - "when future already satisfied" in { + test(s"Promise ($name) $label when future already satisfied") { implicit val timer: Timer = new NullTimer val p = new Promise[Int] p.setValue(3) - assert(p.raiseWithin(1.minute) == p) + assert(use(p, 1.minute, timer) == p) } - "interruption" in Time.withCurrentTimeFrozen { _ => - implicit val timer: Timer = new MockTimer - val p = new HandledPromise[Int] - val f = p.raiseWithin(50.milliseconds) - assert(p.handled.isEmpty) - f.raise(new Exception) - p.handled match { - case Some(_) => - case None => fail() + test(s"Promise ($name) $label interruption") { + Time.withCurrentTimeFrozen { _ => + implicit val timer: Timer = new MockTimer + val p = new HandledPromise[Int] + val f = use(p, 50.milliseconds, timer) + assert(p.handled.isEmpty) + f.raise(new Exception) + p.handled match { + case Some(_) => + case None => fail() + } } } - } - - "masked" should { - "do unconditional interruption" in { - val p = new HandledPromise[Unit] - val f = p.masked - f.raise(new Exception()) - assert(p.handled.isEmpty) - } + } - "do conditional interruption" in { - val p = new HandledPromise[Unit] - val f1 = p.mask { - case _: TimeoutException => true - } - val f2 = p.mask { - case _: TimeoutException => true - } - f1.raise(new TimeoutException("bang!")) - assert(p.handled.isEmpty) - f2.raise(new Exception()) - assert(p.handled.isDefined) - } + test(s"Promise ($name) raiseWithin when we run out of time") { + implicit val timer: Timer = new JavaTimer + val p = new HandledPromise[Int] + intercept[TimeoutException] { + await(p.raiseWithin(50.milliseconds)) } + timer.stop() + p.handled match { + case Some(_) => + case None => fail() + } + } - "liftToTry" should { - "success" in { - val p = const(Return(3)) - assert(await(p.liftToTry) == Return(3)) - } - - "failure" in { - val ex = new Exception() - val p = const(Throw(ex)) - assert(await(p.liftToTry) == Throw(ex)) - } + test(s"Promise ($name) raiseWithin when we run out of time, throw our stuff") { + implicit val timer: Timer = new JavaTimer + class SkyFallException extends Exception("let the skyfall") + val skyFall = new SkyFallException + val p = new HandledPromise[Int] + intercept[SkyFallException] { + await(p.raiseWithin(50.milliseconds, skyFall)) + } + timer.stop() + p.handled match { + case Some(_) => + case None => fail() + } + assert(p.handled.contains(skyFall)) + } - "propagates interrupt" in { - val p = new HandledPromise[Unit] - p.liftToTry.raise(new Exception()) - assert(p.handled.isDefined) - } + test( + s"Promise ($name) raiseWithin when we are within timeout, but inner throws TimeoutException, we don't raise") { + implicit val timer: Timer = new JavaTimer + class SkyFallException extends Exception("let the skyfall") + val skyFall = new SkyFallException + val p = new HandledPromise[Int] + intercept[TimeoutException] { + await( + p.within(20.milliseconds).raiseWithin(50.milliseconds, skyFall) + ) } + timer.stop() + assert(p.handled.isEmpty) + } - "lowerFromTry" should { - "success" in { - val f = const(Return(Return(3))) - assert(await(f.lowerFromTry) == 3) - } + test(s"Promise ($name) raiseWithin when everything is chill") { + implicit val timer: Timer = new JavaTimer + val p = new Promise[Int] + p.setValue(1) + assert(await(p.raiseWithin(50.milliseconds)) == 1) + timer.stop() + } - "failure" in { - val ex = new Exception() - val p = const(Return(Throw(ex))) - val ex1 = intercept[Exception] { await(p.lowerFromTry) } - assert(ex == ex1) - } + test(s"Promise ($name) raiseWithin when timeout is forever") { + // We manage to throw an exception inside + // the scala compiler if we use MockTimer + // here. Sigh. + implicit val timer: Timer = FailingTimer + val p = new Promise[Int] + assert(p.raiseWithin(Duration.Top) == p) + } + + test(s"Promise ($name) raiseWithin when future already satisfied") { + implicit val timer: Timer = new NullTimer + val p = new Promise[Int] + p.setValue(3) + assert(p.raiseWithin(1.minute) == p) + } - "propagates interrupt" in { - val p = new HandledPromise[Try[Unit]] - p.lowerFromTry.raise(new Exception()) - assert(p.handled.isDefined) + test(s"Promise ($name) raiseWithin interruption") { + Time.withCurrentTimeFrozen { _ => + implicit val timer: Timer = new MockTimer + val p = new HandledPromise[Int] + val f = p.raiseWithin(50.milliseconds) + assert(p.handled.isEmpty) + f.raise(new Exception) + p.handled match { + case Some(_) => + case None => fail() } } } - s"FutureTask ($name)" should { - "return result" in { - val task = new FutureTask("hello") - task.run() - assert(await(task) == "hello") - } + test(s"Promise ($name) masked should do unconditional interruption") { + val p = new HandledPromise[Unit] + val f = p.masked + f.raise(new Exception()) + assert(p.handled.isEmpty) + } - "throw result" in { - val task = new FutureTask[String](throw new IllegalStateException) - task.run() - intercept[IllegalStateException] { - await(task) - } + test(s"Promise ($name) masked should do conditional interruption") { + val p = new HandledPromise[Unit] + val f1 = p.mask { + case _: TimeoutException => true } + val f2 = p.mask { + case _: TimeoutException => true + } + f1.raise(new TimeoutException("bang!")) + assert(p.handled.isEmpty) + f2.raise(new Exception()) + assert(p.handled.isDefined) } - } - test("ConstFuture", new MkConst { - def apply[A](r: Try[A]): Future[A] = Future.const(r) - }) - test("Promise", new MkConst { - def apply[A](r: Try[A]): Future[A] = new Promise(r) - }) - - "Future.apply" should { - "fail on NLRC" in { - def ok(): String = { - val f = Future(return "OK") - val t = intercept[FutureNonLocalReturnControl] { - f.poll.get.get() - } - val nlrc = intercept[NonLocalReturnControl[String]] { - throw t.getCause - } - assert(nlrc.value == "OK") - "NOK" - } - assert(ok() == "NOK") + test(s"Promise ($name) liftToTry success") { + val p = const(Return(3)) + assert(await(p.liftToTry) == Return(3)) } - } - "Future.None" should { - "always be defined" in { - assert(Future.None.isDefined) + test(s"Promise ($name) liftToTry failure") { + val ex = new Exception() + val p = const(Throw(ex)) + assert(await(p.liftToTry) == Throw(ex)) } - "but still None" in { - assert(await(Future.None).isEmpty) + + test(s"Promise ($name) liftToTry propagates interrupt") { + val p = new HandledPromise[Unit] + p.liftToTry.raise(new Exception()) + assert(p.handled.isDefined) } - } - "Future.True" should { - "always be defined" in { - assert(Future.True.isDefined) + test(s"Promise ($name) lowerFromTry success") { + val f = const(Return(Return(3))) + assert(await(f.lowerFromTry) == 3) } - "but still True" in { - assert(await(Future.True)) + + test(s"Promise ($name) lowerFromTry failure") { + val ex = new Exception() + val p = const(Return(Throw(ex))) + val ex1 = intercept[Exception] { await(p.lowerFromTry) } + assert(ex == ex1) } - } - "Future.False" should { - "always be defined" in { - assert(Future.False.isDefined) + test(s"Promise ($name) lowerFromTry propagates interrupt") { + val p = new HandledPromise[Try[Unit]] + p.lowerFromTry.raise(new Exception()) + assert(p.handled.isDefined) } - "but still False" in { - assert(!await(Future.False)) + + test(s"FutureTask ($name) should return result") { + val task = new FutureTask("hello") + task.run() + assert(await(task) == "hello") + } + + test(s"FutureTask ($name) should throw result") { + val task = new FutureTask[String](throw new IllegalStateException) + task.run() + intercept[IllegalStateException] { + await(task) + } } } - "Future.never" should { - "must be undefined" in { - assert(!Future.never.isDefined) - assert(Future.never.poll.isEmpty) + tests( + "ConstFuture", + new MkConst { + def apply[A](r: Try[A]): Future[A] = Future.const(r) + }) + tests( + "Promise", + new MkConst { + def apply[A](r: Try[A]): Future[A] = new Promise(r) + }) + tests( + "Future.fromCompletableFuture", + new MkConst { + def apply[A](r: Try[A]): Future[A] = + Future.fromCompletableFuture(new Promise(r).toCompletableFuture[A]) } + ) - "always time out" in { - intercept[TimeoutException] { Await.ready(Future.never, 0.milliseconds) } + test("Future.apply should fail on NLRC") { + def ok(): String = { + val f = Future(return "OK") + val t = intercept[FutureNonLocalReturnControl] { + f.poll.get.get() + } + val nlrc = intercept[NonLocalReturnControl[String]] { + throw t.getCause + } + assert(nlrc.value == "OK") + "NOK" } + assert(ok() == "NOK") + } + + test("Future.None should always be defined") { + assert(Future.None.isDefined) + } + test("Future.None should always be definedbut still None") { + assert(await(Future.None).isEmpty) + } + + test("Future.True should always be defined") { + assert(Future.True.isDefined) + } + test("Future.True should always be definedbut still True") { + assert(await(Future.True)) + } + + test("Future.False should always be defined") { + assert(Future.False.isDefined) + } + test("Future.False should always be defined but still False") { + assert(!await(Future.False)) + } + + test("Future.never must be undefined") { + assert(!Future.never.isDefined) + assert(Future.never.poll.isEmpty) + } + + test("Future.never should always time out") { + intercept[TimeoutException] { Await.ready(Future.never, 0.milliseconds) } } - "Future.onFailure" should { + test("Future.onFailure with Function1") { val nonfatal = Future.exception(new RuntimeException()) val fatal = Future.exception(new FatalException()) + val counter = new AtomicInteger() + val f: Throwable => Unit = _ => counter.incrementAndGet() + nonfatal.onFailure(f) + assert(counter.get() == 1) + fatal.onFailure(f) + assert(counter.get() == 2) + } - "with Function1" in { + test("Future.onFailure with PartialFunction") { + val nonfatal = Future.exception(new RuntimeException()) + val fatal = Future.exception(new FatalException()) + val monitor = new HandledMonitor() + Monitor.using(monitor) { val counter = new AtomicInteger() - val f: Throwable => Unit = _ => counter.incrementAndGet() - nonfatal.onFailure(f) + nonfatal.onFailure { case NonFatal(_) => counter.incrementAndGet() } assert(counter.get() == 1) - fatal.onFailure(f) - assert(counter.get() == 2) - } - - "with PartialFunction" in { - val monitor = new HandledMonitor() - Monitor.using(monitor) { - val counter = new AtomicInteger() - nonfatal.onFailure { case NonFatal(_) => counter.incrementAndGet() } - assert(counter.get() == 1) - assert(monitor.handled == null) + assert(monitor.handled == null) - // this will throw a MatchError and propagated to the monitor - fatal.onFailure { case NonFatal(_) => counter.incrementAndGet() } - assert(counter.get() == 1) - assert(monitor.handled.getClass == classOf[MatchError]) - } + // this will throw a MatchError and propagated to the monitor + fatal.onFailure { case NonFatal(_) => counter.incrementAndGet() } + assert(counter.get() == 1) + assert(monitor.handled.getClass == classOf[MatchError]) } } - "Future.sleep" should { - "Satisfy after the given amount of time" in Time.withCurrentTimeFrozen { tc => + test("Future.sleep should Satisfy after the given amount of time") { + Time.withCurrentTimeFrozen { tc => implicit val timer: MockTimer = new MockTimer val f = Future.sleep(10.seconds) @@ -1768,217 +1801,204 @@ class FutureTest extends WordSpec with MockitoSugar with GeneratorDrivenProperty assert(f.isDefined) await(f) } + } - "Be interruptible" in { - implicit val timer: MockTimer = new MockTimer + test("Future.sleep should Be interruptible") { + implicit val timer: MockTimer = new MockTimer - // sleep and grab the task that's created - val f = Future.sleep(1.second)(timer) - val task = timer.tasks(0) - // then raise a known exception - val e = new Exception("expected") - f.raise(e) + // sleep and grab the task that's created + val f = Future.sleep(1.second)(timer) + val task = timer.tasks(0) + // then raise a known exception + val e = new Exception("expected") + f.raise(e) - // we were immediately satisfied with the exception and the task was canceled - f mustProduce Throw(e) - assert(task.isCancelled) - } + // we were immediately satisfied with the exception and the task was canceled + f mustProduce Throw(e) + assert(task.isCancelled) + } - "Return Future.Done for durations <= 0" in { - implicit val timer: MockTimer = new MockTimer - assert(Future.sleep(Duration.Zero) eq Future.Done) - assert(Future.sleep((-10).seconds) eq Future.Done) - assert(timer.tasks.isEmpty) - } + test("Future.sleep should Return Future.Done for durations <= 0") { + implicit val timer: MockTimer = new MockTimer + assert(Future.sleep(Duration.Zero) eq Future.Done) + assert(Future.sleep((-10).seconds) eq Future.Done) + assert(timer.tasks.isEmpty) + } - "Return Future.never for Duration.Top" in { - implicit val timer: MockTimer = new MockTimer - assert(Future.sleep(Duration.Top) eq Future.never) - assert(timer.tasks.isEmpty) - } + test("Future.sleep should Return Future.never for Duration.Top") { + implicit val timer: MockTimer = new MockTimer + assert(Future.sleep(Duration.Top) eq Future.never) + assert(timer.tasks.isEmpty) } - "Future.select" should { - import Arbitrary.arbitrary + test("Future.select should return the first result") { val genLen = Gen.choose(1, 10) - - "return the first result" in { - forAll(genLen, arbitrary[Boolean]) { (n, fail) => - val ps = List.fill(n)(new Promise[Int]()) - assert(ps.map(_.waitqLength).sum == 0) - val f = Future.select(ps) - val i = Random.nextInt(ps.length) - val e = new Exception("sad panda") - val t = if (fail) Throw(e) else Return(i) - ps(i).update(t) - assert(f.isDefined) - val (ft, fps) = await(f) - assert(ft == t) - assert(fps.toSet == (ps.toSet - ps(i))) - } - } - - "not accumulate listeners when losing or" in { - val p = new Promise[Unit] - val q = new Promise[Unit] - p.or(q) - assert(p.waitqLength == 1) - q.setDone() - assert(p.waitqLength == 0) - } - - "not accumulate listeners when losing select" in { - val p = new Promise[Unit] - val q = new Promise[Unit] - Future.select(Seq(p, q)) - assert(p.waitqLength == 1) - q.setDone() - assert(p.waitqLength == 0) - } - - "not accumulate listeners if not selected" in { - forAll(genLen, arbitrary[Boolean]) { (n, fail) => - val ps = List.fill(n)(new Promise[Int]()) - assert(ps.map(_.waitqLength).sum == 0) - val f = Future.select(ps) - assert(ps.map(_.waitqLength).sum == n) - val i = Random.nextInt(ps.length) - val e = new Exception("sad panda") - val t = if (fail) Throw(e) else Return(i) - f.respond { _ => - () - } - assert(ps.map(_.waitqLength).sum == n) - ps(i).update(t) - assert(ps.map(_.waitqLength).sum == 0) - } - } - - "fail if we attempt to select an empty future sequence" in { - val f = Future.select(Nil) + forAll(genLen, arbitrary[Boolean]) { (n, fail) => + val ps = List.fill(n)(new Promise[Int]()) + assert(ps.map(_.waitqLength).sum == 0) + val f = Future.select(ps) + val i = Random.nextInt(ps.length) + val e = new Exception("sad panda") + val t = if (fail) Throw(e) else Return(i) + ps(i).update(t) assert(f.isDefined) - val e = new IllegalArgumentException("empty future list") - val actual = intercept[IllegalArgumentException] { await(f) } - assert(actual.getMessage == e.getMessage) + val (ft, fps) = await(f) + assert(ft == t) + assert(fps.toSet == (ps.toSet - ps(i))) } + } + + test("Future.select should not accumulate listeners when losing or") { + val p = new Promise[Unit] + val q = new Promise[Unit] + p.or(q) + assert(p.waitqLength == 1) + q.setDone() + assert(p.waitqLength == 0) + } + + test("Future.select should not accumulate listeners when losing select") { + val p = new Promise[Unit] + val q = new Promise[Unit] + Future.select(Seq(p, q)) + assert(p.waitqLength == 1) + q.setDone() + assert(p.waitqLength == 0) + } - "propagate interrupts" in { - val fs = (0 until 10).map(_ => new HandledPromise[Int]) - Future.select(fs).raise(new Exception) - assert(fs.forall(_.handled.isDefined)) + test("Future.select should not accumulate listeners if not selected") { + val genLen = Gen.choose(1, 10) + forAll(genLen, arbitrary[Boolean]) { (n, fail) => + val ps = List.fill(n)(new Promise[Int]()) + assert(ps.map(_.waitqLength).sum == 0) + val f = Future.select(ps) + assert(ps.map(_.waitqLength).sum == n) + val i = Random.nextInt(ps.length) + val e = new Exception("sad panda") + val t = if (fail) Throw(e) else Return(i) + f.respond { _ => () } + assert(ps.map(_.waitqLength).sum == n) + ps(i).update(t) + assert(ps.map(_.waitqLength).sum == 0) } } - // These tests are almost a carbon copy of the "Future.select" tests, they + test("Future.select should fail if we attempt to select an empty future sequence") { + val f = Future.select(Nil) + assert(f.isDefined) + val e = new IllegalArgumentException("empty future list") + val actual = intercept[IllegalArgumentException] { await(f) } + assert(actual.getMessage == e.getMessage) + } + + test("Future.select should propagate interrupts") { + val fs = (0 until 10).map(_ => new HandledPromise[Int]) + Future.select(fs).raise(new Exception) + assert(fs.forall(_.handled.isDefined)) + } + + // These "Future.selectIndex" tests are almost a carbon copy of the "Future.select" tests, they // should evolve in-sync. - "Future.selectIndex" should { - import Arbitrary.arbitrary + test("Future.selectIndex should return the first result") { val genLen = Gen.choose(1, 10) - "return the first result" in { - forAll(genLen, arbitrary[Boolean]) { (n, fail) => - val ps = IndexedSeq.fill(n)(new Promise[Int]()) - assert(ps.map(_.waitqLength).sum == 0) - val f = Future.selectIndex(ps) - val i = Random.nextInt(ps.length) - val e = new Exception("sad panda") - val t = if (fail) Throw(e) else Return(i) - ps(i).update(t) - assert(f.isDefined) - assert(await(f) == i) - } - } - - "not accumulate listeners when losing select" in { - val p = new Promise[Unit] - val q = new Promise[Unit] - Future.selectIndex(IndexedSeq(p, q)) - assert(p.waitqLength == 1) - q.setDone() - assert(p.waitqLength == 0) - } - - "not accumulate listeners if not selected" in { - forAll(genLen, arbitrary[Boolean]) { (n, fail) => - val ps = IndexedSeq.fill(n)(new Promise[Int]()) - assert(ps.map(_.waitqLength).sum == 0) - val f = Future.selectIndex(ps) - assert(ps.map(_.waitqLength).sum == n) - val i = Random.nextInt(ps.length) - val e = new Exception("sad panda") - val t = if (fail) Throw(e) else Return(i) - f.respond { _ => - () - } - assert(ps.map(_.waitqLength).sum == n) - ps(i).update(t) - assert(ps.map(_.waitqLength).sum == 0) - } - } - - "fail if we attempt to select an empty future sequence" in { - val f = Future.selectIndex(IndexedSeq.empty) + forAll(genLen, arbitrary[Boolean]) { (n, fail) => + val ps = IndexedSeq.fill(n)(new Promise[Int]()) + assert(ps.map(_.waitqLength).sum == 0) + val f = Future.selectIndex(ps) + val i = Random.nextInt(ps.length) + val e = new Exception("sad panda") + val t = if (fail) Throw(e) else Return(i) + ps(i).update(t) assert(f.isDefined) - val e = new IllegalArgumentException("empty future list") - val actual = intercept[IllegalArgumentException] { await(f) } - assert(actual.getMessage == e.getMessage) + assert(await(f) == i) } + } - "propagate interrupts" in { - val fs = IndexedSeq.fill(10)(new HandledPromise[Int]()) - Future.selectIndex(fs).raise(new Exception) - assert(fs.forall(_.handled.isDefined)) + test("Future.selectIndex should not accumulate listeners when losing select") { + val p = new Promise[Unit] + val q = new Promise[Unit] + Future.selectIndex(IndexedSeq(p, q)) + assert(p.waitqLength == 1) + q.setDone() + assert(p.waitqLength == 0) + } + + test("Future.selectIndex should not accumulate listeners if not selected") { + val genLen = Gen.choose(1, 10) + forAll(genLen, arbitrary[Boolean]) { (n, fail) => + val ps = IndexedSeq.fill(n)(new Promise[Int]()) + assert(ps.map(_.waitqLength).sum == 0) + val f = Future.selectIndex(ps) + assert(ps.map(_.waitqLength).sum == n) + val i = Random.nextInt(ps.length) + val e = new Exception("sad panda") + val t = if (fail) Throw(e) else Return(i) + f.respond { _ => () } + assert(ps.map(_.waitqLength).sum == n) + ps(i).update(t) + assert(ps.map(_.waitqLength).sum == 0) } } - "Future.each" should { - "iterate until an exception is thrown" in { - val exc = new Exception("done") - var next: Future[Int] = Future.value(10) - val done = Future.each(next) { - case 0 => next = Future.exception(exc) - case n => next = Future.value(n - 1) - } + test("Future.selectIndex should fail if we attempt to select an empty future sequence") { + val f = Future.selectIndex(IndexedSeq.empty) + assert(f.isDefined) + val e = new IllegalArgumentException("empty future list") + val actual = intercept[IllegalArgumentException] { await(f) } + assert(actual.getMessage == e.getMessage) + } - assert(done.poll.contains(Throw(exc))) + test("Future.selectIndex should propagate interrupts") { + val fs = IndexedSeq.fill(10)(new HandledPromise[Int]()) + Future.selectIndex(fs).raise(new Exception) + assert(fs.forall(_.handled.isDefined)) + } + + test("Future.each should iterate until an exception is thrown") { + val exc = new Exception("done") + var next: Future[Int] = Future.value(10) + val done = Future.each(next) { + case 0 => next = Future.exception(exc) + case n => next = Future.value(n - 1) } - "evaluate next one time per iteration" in { - var i, j = 0 - def next(): Future[Int] = - if (i == 10) Future.exception(new Exception) - else { - i += 1 - Future.value(i) - } + assert(done.poll.contains(Throw(exc))) + } - Future.each(next()) { i => - j += 1 - assert(i == j) + test("Future.each should evaluate next one time per iteration") { + var i, j = 0 + def next(): Future[Int] = + if (i == 10) Future.exception(new Exception) + else { + i += 1 + Future.value(i) } - } - "terminate if the body throws an exception" in { - val exc = new Exception("body exception") - var i = 0 - def next(): Future[Int] = Future.value({ i += 1; i }) - val done = Future.each(next()) { - case 10 => throw exc - case _ => - } + Future.each(next()) { i => + j += 1 + assert(i == j) + } + } - assert(done.poll.contains(Throw(exc))) - assert(i == 10) + test("Future.each should terminate if the body throws an exception") { + val exc = new Exception("body exception") + var i = 0 + def next(): Future[Int] = Future.value({ i += 1; i }) + val done = Future.each(next()) { + case 10 => throw exc + case _ => } - "terminate when 'next' throws" in { - val exc = new Exception - def next(): Future[Int] = throw exc - val done = Future.each(next()) { _ => - throw exc - } + assert(done.poll.contains(Throw(exc))) + assert(i == 10) + } - assert(done.poll.contains(Throw(Future.NextThrewException(exc)))) - } + test("Future.each should terminate when 'next' throws") { + val exc = new Exception + def next(): Future[Int] = throw exc + val done = Future.each(next()) { _ => throw exc } + + assert(done.poll.contains(Throw(Future.NextThrewException(exc)))) } } diff --git a/util-core/src/test/scala/com/twitter/util/ImmediateValueFutureTest.scala b/util-core/src/test/scala/com/twitter/util/ImmediateValueFutureTest.scala new file mode 100644 index 0000000000..bc650e7765 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/util/ImmediateValueFutureTest.scala @@ -0,0 +1,221 @@ +package com.twitter.util + +import com.twitter.concurrent.Scheduler +import java.util.concurrent.atomic.AtomicBoolean +import org.scalatest.funsuite.AnyFunSuite + +class ImmediateValueFutureTest extends AnyFunSuite { + + // The following is intended to demonstrate the differing recursive behaviour of `ConstFuture` + // and `ImmediateValueFuture`. `ConstFuture` essentially does tail-call elimination which means + // the stack will not grow during a recursive call. `ImmediateValueFuture` does not, so the stack + // will grow with each call! + def recurseAndGetStackSizes(f: Future[Unit]): Seq[Int] = { + val stop = new AtomicBoolean(false) + + @volatile var loopStackSizes: Seq[Int] = Seq.empty + + def loop(): Future[Unit] = { + if (stop.get) { + Future.Done + } else { + f.flatMap { _ => + loopStackSizes = loopStackSizes :+ new Throwable().getStackTrace.length + loop() + } + } + } + + FuturePool.unboundedPool { + loop() + } + + while (loopStackSizes.size < 10) {} + stop.set(true) + loopStackSizes + } + + test("ConstFuture recursion does not grow the stack") { + val loopStackSizes = recurseAndGetStackSizes(Future.const(Return[Unit](()))) + assert(loopStackSizes.forall(_ == loopStackSizes.head)) + } + + test("ImmediateValueFuture recursion does grow the stack") { + val loopStackSizes = recurseAndGetStackSizes(new ImmediateValueFuture(())) + assert(loopStackSizes.zip(loopStackSizes.tail).forall { case (a, b) => a < b }) + } + + test(s"ImmediateValueFuture.interruptible should do nothing") { + val f = new ImmediateValueFuture(()) + val i = f.interruptible() + i.raise(new Exception()) + assert(f.poll.contains(Return[Unit](()))) + assert(i.poll.contains(Return[Unit](()))) + } + + test(s"ImmediateValueFuture should propagate locals and restore original context in `respond`") { + val local = new Local[Int] + val f = new ImmediateValueFuture(111) + + var ran = 0 + local() = 1010 + + f.ensure { + assert(local().contains(1010)) + local() = 1212 + f.ensure { + assert(local().contains(1212)) + local() = 1313 + ran += 1 + } + assert(local().contains(1212)) + ran += 1 + } + + assert(local().contains(1010)) + assert(ran == 2) + } + + test( + s"ImmediateValueFuture should propagate locals and restore original context in `transform`") { + val local = new Local[Int] + val f = new ImmediateValueFuture(111) + + var ran = 0 + local() = 1010 + + f.transform { tryRes => + assert(local().contains(1010)) + local() = 1212 + f.transform { tryRes => + assert(local().contains(1212)) + local() = 1313 + ran += 1 + Future.const(tryRes) + } + assert(local().contains(1212)) + ran += 1 + Future.const(tryRes) + } + + assert(local().contains(1010)) + assert(ran == 2) + } + + test(s"ImmediateValueFuture should propagate locals and restore original context in `map`") { + val local = new Local[Int] + val f = new ImmediateValueFuture(111) + + var ran = 0 + local() = 1010 + + f.map { i => + assert(local().contains(1010)) + local() = 1212 + f.map { i => + assert(local().contains(1212)) + local() = 1313 + ran += 1 + i + } + assert(local().contains(1212)) + ran += 1 + i + } + + assert(local().contains(1010)) + assert(ran == 2) + } + + test(s"ImmediateValueFuture should not delay execution") { + val numDispatchesBefore = Scheduler().numDispatches + val f = new ImmediateValueFuture(111) + + var count = 0 + f.onSuccess { _ => + assert(count == 0) + f.ensure { + assert(count == 0) + count += 1 + } + + assert(count == 1) + count += 1 + } + + assert(count == 2) + assert(Scheduler().numDispatches == numDispatchesBefore) + } + + test(s"ImmediateValueFuture side effects should be monitored") { + val inner = new ImmediateValueFuture(111) + val exc = new Exception("a raw exception") + + var monitored = false + + val monitor = new Monitor { + override def handle(exc: Throwable): Boolean = { + monitored = true + true + } + } + + Monitor.using(monitor) { + val f = inner.respond { _ => + throw exc + } + + assert(f.poll.contains(Return(111))) + assert(monitored == true) + } + } + + test("ImmediateValueFuture.rescue returns self") { + val f = new ImmediateValueFuture(111) + val r = f.rescue { + case e: Exception => Future.value(1) + } + + assert(f eq r) + } + + test("ImmediateValueFuture.map returns ImmediateValueFuture if f does not throw") { + val f1 = new ImmediateValueFuture(111) + + val f2 = f1.map { _ => + "hello" + } + + assert(f2.isInstanceOf[ImmediateValueFuture[String]]) + } + + test("ImmediateValueFuture.map returns Future.exception if f throws") { + val f1 = new ImmediateValueFuture(111) + + val f2 = f1.map { _ => + throw new Exception("boom!") + } + + intercept[Exception](Await.result(f2)) + } + + test("ImmediateValueFuture.flatMap returns Future.exception if f returns exceptional Future") { + val f1 = new ImmediateValueFuture(111) + + val f2 = f1.flatMap { _ => + Future.exception(new Exception("boom!")) + } + + intercept[Exception](Await.result(f2)) + } + + test("ImmediateValueFuture.transform returns Future.exception if f throws") { + val f1 = new ImmediateValueFuture(111) + + val f2 = f1.transform { _ => + throw new Exception("boom!") + } + + intercept[Exception](Await.result(f2)) + } +} diff --git a/util-core/src/test/scala/com/twitter/util/LastWriteWinsQueueTest.scala b/util-core/src/test/scala/com/twitter/util/LastWriteWinsQueueTest.scala index a799d9be8f..da989d1fc9 100644 --- a/util-core/src/test/scala/com/twitter/util/LastWriteWinsQueueTest.scala +++ b/util-core/src/test/scala/com/twitter/util/LastWriteWinsQueueTest.scala @@ -1,29 +1,24 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class LastWriteWinsQueueTest extends WordSpec { - "LastWriteWinsQueue" should { - val queue = new LastWriteWinsQueue[String] +class LastWriteWinsQueueTest extends AnyFunSuite { + val queue = new LastWriteWinsQueue[String] - "add & remove items" in { - assert(queue.size == 0) - queue.add("1") - assert(queue.size == 1) - assert(queue.remove() == "1") - assert(queue.size == 0) - } + test("LastWriteWinsQueue should add & remove items") { + assert(queue.size == 0) + queue.add("1") + assert(queue.size == 1) + assert(queue.remove == "1") + assert(queue.size == 0) + } - "last write wins" in { - queue.add("1") - queue.add("2") - assert(queue.size == 1) - assert(queue.poll() == "2") - assert(queue.size == 0) - assert(queue.poll() == null) - } + test("LastWriteWinsQueue should last write wins") { + queue.add("1") + queue.add("2") + assert(queue.size == 1) + assert(queue.poll() == "2") + assert(queue.size == 0) + assert(queue.poll() == null) } } diff --git a/util-core/src/test/scala/com/twitter/util/LocalTest.scala b/util-core/src/test/scala/com/twitter/util/LocalTest.scala index 9651b6e2f9..ee925fe1fd 100644 --- a/util-core/src/test/scala/com/twitter/util/LocalTest.scala +++ b/util-core/src/test/scala/com/twitter/util/LocalTest.scala @@ -2,10 +2,10 @@ package com.twitter.util import org.scalacheck.Arbitrary import org.scalacheck.Arbitrary.arbitrary -import org.scalatest.FunSuite -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -class LocalTest extends FunSuite with GeneratorDrivenPropertyChecks { +class LocalTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { test("Local: should be undefined by default") { val local = new Local[Int] assert(local() == None) @@ -129,6 +129,22 @@ class LocalTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(local() == Some(123)) } + test("Local.let: should only restore the given key, value pair") { + val l1, l2 = new Local[Int] + l1() = 1 + l2() = 2 + val ctx = Local.save() + + l1.let(2) { + l2() = 4 + assert(l1() == Some(2)) + assert(l2() == Some(4)) + } + + assert(l1() == Some(1)) + assert(l2() == Some(4)) + } + test("Local.letClear: should clear Local and restore previous value") { val local = new Local[Int] local() = 123 @@ -138,6 +154,23 @@ class LocalTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(local() == Some(123)) } + test("Local.letClear: should only restore the given key, value pair") { + val l1, l2 = new Local[Int] + l1() = 1 + l2() = 2 + val ctx = Local.save() + + l1.letClear { + assert(l1() == None) + assert(l2() == Some(2)) + l2() = 4 + assert(l2() == Some(4)) + } + + assert(l1() == Some(1)) + assert(l2() == Some(4)) + } + test("Local.clear: should make a copy when clearing") { val l = new Local[Int] l() = 1 @@ -163,11 +196,62 @@ class LocalTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(fn() == 101) } + test("Local.threadLocalGetter returns the same value as local.apply") { + val l = new Local[Int] + val getter = l.threadLocalGetter() + l() = 1 + assert(getter() == l()) + val adder: () => Int = { () => + val rv = 100 + getter().get + l() = 10000 + assert(getter() == l()) + rv + } + val fn = Local.closed(adder) + l() = 100 + assert(getter() == l()) + assert(fn() == 101) + assert(getter() == l()) + assert(l() == Some(100)) + assert(fn() == 101) + } + + test("Local.threadLocalGetter should have a per-thread state") { + val local = new Local[Int] + val getter = local.threadLocalGetter() + var threadValue: Option[Int] = null + + local() = 123 + + val t = new Thread { + override def run() = { + val getter2 = local.threadLocalGetter() + assert(getter2() == None) + local() = 333 + threadValue = getter2() + } + } + + t.start() + t.join() + + assert(local() == Some(123)) + assert(getter() == Some(123)) + assert(threadValue == Some(333)) + } + + test("Local.threadLocalGetter asserts that it used by its owner") { + val l = new Local[Int] + val getter = + Await.result(FuturePool.unboundedPool(l.threadLocalGetter())) + intercept[AssertionError](getter()) + } + implicit lazy val arbKey: Arbitrary[Local.Key] = Arbitrary(new Local.Key) test("Context should have the same behavior with Map") { val context = Local.Context.empty - val map = Map.empty[Local.Key, Option[Int]] + val map = Map.empty[Local.Key, Int] forAll(arbitrary[Int], arbitrary[Local.Key]) { (value, key) => if (value % 2 == 0) { diff --git a/util-core/src/test/scala/com/twitter/util/MemoizeTest.scala b/util-core/src/test/scala/com/twitter/util/MemoizeTest.scala index bc7ae18741..1e72509627 100644 --- a/util-core/src/test/scala/com/twitter/util/MemoizeTest.scala +++ b/util-core/src/test/scala/com/twitter/util/MemoizeTest.scala @@ -1,15 +1,14 @@ package com.twitter.util -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import java.util.concurrent.atomic.AtomicInteger -import java.util.concurrent.{TimeUnit, CountDownLatch => JavaCountDownLatch} -import org.junit.runner.RunWith +import java.util.concurrent.atomic.AtomicReference +import java.util.concurrent.{CountDownLatch => JavaCountDownLatch} +import java.util.concurrent.TimeUnit import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class MemoizeTest extends FunSuite { +class MemoizeTest extends AnyFunSuite { test("Memoize.apply: only runs the function once for the same input") { // mockito can't spy anonymous classes, // and this was the simplest approach i could come up with. @@ -50,7 +49,7 @@ class MemoizeTest extends FunSuite { val callCount = new AtomicInteger(0) val startUpLatch = new JavaCountDownLatch(1) - val memoizer = Memoize { i: Int => + val memoizer = Memoize { (i: Int) => // Wait for all of the threads to be started before // continuing. This gives races a chance to happen. startUpLatch.await() @@ -66,9 +65,7 @@ class MemoizeTest extends FunSuite { val ConcurrencyLevel = 5 val computations = - Future.collect(1 to ConcurrencyLevel map { _ => - FuturePool.unboundedPool(memoizer(5)) - }) + Future.collect(1 to ConcurrencyLevel map { _ => FuturePool.unboundedPool(memoizer(5)) }) startUpLatch.countDown() val results = Await.result(computations) @@ -90,7 +87,7 @@ class MemoizeTest extends FunSuite { // A computation that should fail the first time, and then // succeed for all subsequent attempts. - val memo = Memoize { i: Int => + val memo = Memoize { (i: Int) => // Ensure that all of the callers have been started startUpLatch.await(200, TimeUnit.MILLISECONDS) // This effect should happen once per exception plus once for @@ -139,6 +136,7 @@ class MemoizeTest extends FunSuite { assert(2 == memoizer(1)) assert(3 == memoizer(2)) assertResult(Map(1 -> 2, 2 -> 3))(memoizer.snap) + assert(memoizer.size == 2) } test("Memoize.snappable: snap ignores in-process computations") { @@ -164,7 +162,10 @@ class MemoizeTest extends FunSuite { callReadyLatch.countDown() assert(3 == Await.result(result)) - assertResult(Map(1 -> 2, 2 -> 3))(memoizer.snap) + val snap = memoizer.snap + assertResult(Map(1 -> 2, 2 -> 3))(snap) + assert(memoizer.size == snap.size) + assert(snap.size == 2) } test("Memoize.snappable: snap ignores failed computations") { @@ -178,8 +179,54 @@ class MemoizeTest extends FunSuite { memoizer(2) } assert(memoizer.snap.isEmpty) + assert(memoizer.size == memoizer.snap.size) + assert(memoizer.snap.size == 0) assert(2 == memoizer(1)) assertResult(Map(1 -> 2))(memoizer.snap) + assert(memoizer.size == memoizer.snap.size) + assert(memoizer.snap.size == 1) + } + + test("Memoize.classValue: returns the same value for the same input") { + val cv = Memoize.classValue { _ => + new Object() + } + // just some anonymous classes + val clazz1 = new Object() {}.getClass + val clazz2 = new Object() {}.getClass + assert(clazz1 ne clazz2) + + val o11 = cv(clazz1) + val o12 = cv(clazz1) + val o21 = cv(clazz2) + val o22 = cv(clazz2) + assert(o11 eq o12) + assert(o21 eq o22) + assert(o11 ne o21) + } + + test("Memoize.classValue: racing threads observe the same value") { + val Concurrency = 16 + val ready = new JavaCountDownLatch(Concurrency) + val cv = Memoize.classValue { _ => + ready.await(1, TimeUnit.SECONDS) + new Object() + } + + val clazz = new Object() {}.getClass + + val values = new Array[AtomicReference[Object]](Concurrency) + val threads = (0 until Concurrency).map { i => + new Thread(() => { + ready.countDown() + values(i) = new AtomicReference(cv(clazz)) + }) + } + threads.foreach { _.start() } + threads.foreach { _.join(1000L) } + + val v0 = values(0).get() + assert(values.forall(_.get eq v0)) } } diff --git a/util-core/src/test/scala/com/twitter/util/MonitorTest.scala b/util-core/src/test/scala/com/twitter/util/MonitorTest.scala index b48d4025ac..9f17351274 100644 --- a/util-core/src/test/scala/com/twitter/util/MonitorTest.scala +++ b/util-core/src/test/scala/com/twitter/util/MonitorTest.scala @@ -1,9 +1,9 @@ package com.twitter.util -import org.mockito.Matchers._ +import org.mockito.ArgumentMatchers._ import org.mockito.Mockito._ -import org.scalatest.WordSpec -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite object MonitorTest { class MockMonitor extends Monitor { @@ -11,128 +11,118 @@ object MonitorTest { } } -class MonitorTest extends WordSpec with MockitoSugar { +class MonitorTest extends AnyFunSuite with MockitoSugar { import MonitorTest._ - "Monitor#orElse" should { - class MonitorOrElseHelper { - val m0, m1, m2 = spy(new MockMonitor()) - Seq(m0, m1, m2).foreach { m => - when(m.handle(any[Throwable])).thenReturn(true) - } - val exc = new Exception - val m = m0.orElse(m1).orElse(m2) - } + class MonitorOrElseHelper { + val m0, m1, m2 = spy(new MockMonitor()) + Seq(m0, m1, m2).foreach { m => when(m.handle(any[Throwable])).thenReturn(true) } + val exc = new Exception + val m = m0.orElse(m1).orElse(m2) + } - "stop at first successful handle" in { - val h = new MonitorOrElseHelper - import h._ + test("Monitor#orElse should stop at first successful handle") { + val h = new MonitorOrElseHelper + import h._ - assert(m.handle(exc)) + assert(m.handle(exc)) - verify(m0).handle(exc) - verify(m1, never()).handle(exc) - verify(m2, never()).handle(exc) + verify(m0).handle(exc) + verify(m1, never()).handle(exc) + verify(m2, never()).handle(exc) - when(m0.handle(any[Throwable])).thenReturn(false) + when(m0.handle(any[Throwable])).thenReturn(false) - assert(m.handle(exc)) - verify(m0, times(2)).handle(exc) - verify(m1).handle(exc) - verify(m2, never()).handle(exc) - } + assert(m.handle(exc)) + verify(m0, times(2)).handle(exc) + verify(m1).handle(exc) + verify(m2, never()).handle(exc) + } - "fail when no nothing got handled" in { - val h = new MonitorOrElseHelper - import h._ + test("Monitor#orElse should fail when no nothing got handled") { + val h = new MonitorOrElseHelper + import h._ - Seq(m0, m1, m2) foreach { m => - when(m.handle(any[Throwable])).thenReturn(false) - } - assert(!m.handle(exc)) - Seq(m0, m1, m2) foreach { m => - verify(m).handle(exc) - } - } + Seq(m0, m1, m2) foreach { m => when(m.handle(any[Throwable])).thenReturn(false) } + assert(!m.handle(exc)) + Seq(m0, m1, m2) foreach { m => verify(m).handle(exc) } + } - "wrap Monitor exceptions and pass them on" in { - val h = new MonitorOrElseHelper - import h._ + test("Monitor#orElse should wrap Monitor exceptions and pass them on") { + val h = new MonitorOrElseHelper + import h._ - val rte = new RuntimeException("really bad news") - when(m0.handle(any[Throwable])).thenThrow(rte) + val rte = new RuntimeException("really bad news") + when(m0.handle(any[Throwable])).thenThrow(rte) - assert(m.handle(exc)) - verify(m0).handle(exc) - verify(m1).handle(MonitorException(exc, rte)) - } + assert(m.handle(exc)) + verify(m0).handle(exc) + verify(m1).handle(MonitorException(exc, rte)) } - "Monitor#andThen" should { - class MonitorAndThenHelper { - val m0, m1 = spy(new MockMonitor) - when(m0.handle(any[Throwable])).thenReturn(true) - when(m1.handle(any[Throwable])).thenReturn(true) - val m = m0.andThen(m1) - val exc = new Exception - } + class MonitorAndThenHelper { + val m0, m1 = spy(new MockMonitor) + when(m0.handle(any[Throwable])).thenReturn(true) + when(m1.handle(any[Throwable])).thenReturn(true) + val m = m0.andThen(m1) + val exc = new Exception + } - "run all monitors" in { - val h = new MonitorAndThenHelper - import h._ + test("Monitor#andThen should run all monitors") { + val h = new MonitorAndThenHelper + import h._ - when(m0.handle(any[Throwable])).thenReturn(true) - when(m1.handle(any[Throwable])).thenReturn(true) + when(m0.handle(any[Throwable])).thenReturn(true) + when(m1.handle(any[Throwable])).thenReturn(true) - assert(m.handle(exc)) - verify(m0).handle(exc) - verify(m1).handle(exc) - } + assert(m.handle(exc)) + verify(m0).handle(exc) + verify(m1).handle(exc) + } - "be succcessful when any underlying monitor is" in { - val h = new MonitorAndThenHelper - import h._ + test("Monitor#andThen should be succcessful when any underlying monitor is") { + val h = new MonitorAndThenHelper + import h._ - when(m0.handle(any[Throwable])).thenReturn(false) - assert(m.handle(exc)) - when(m1.handle(any[Throwable])).thenReturn(false) - assert(!m.handle(exc)) - } + when(m0.handle(any[Throwable])).thenReturn(false) + assert(m.handle(exc)) + when(m1.handle(any[Throwable])).thenReturn(false) + assert(!m.handle(exc)) + } - "wrap Monitor exceptions and pass them on" in { - val h = new MonitorAndThenHelper - import h._ + test("Monitor#andThen should wrap Monitor exceptions and pass them on") { + val h = new MonitorAndThenHelper + import h._ - val rte = new RuntimeException("really bad news") - when(m0.handle(any[Throwable])).thenThrow(rte) + val rte = new RuntimeException("really bad news") + when(m0.handle(any[Throwable])).thenThrow(rte) - assert(m.handle(exc)) - verify(m0).handle(exc) - verify(m1).handle(MonitorException(exc, rte)) - } + assert(m.handle(exc)) + verify(m0).handle(exc) + verify(m1).handle(MonitorException(exc, rte)) + } - "fail if both monitors throw" in { - val h = new MonitorAndThenHelper - import h._ + test("Monitor#andThen should fail if both monitors throw") { + val h = new MonitorAndThenHelper + import h._ - val rte = new RuntimeException("really bad news") - when(m0.handle(any[Throwable])).thenThrow(rte) - when(m1.handle(any[Throwable])).thenThrow(rte) + val rte = new RuntimeException("really bad news") + when(m0.handle(any[Throwable])).thenThrow(rte) + when(m1.handle(any[Throwable])).thenThrow(rte) - assert(!m.handle(exc)) - } + assert(!m.handle(exc)) } - "Monitor.get, Monitor.set()" should { + test("Monitor.get, Monitor.set() should maintain current monitor") { val m = new MockMonitor - "maintain current monitor" in Monitor.restoring { + Monitor.restoring { Monitor.set(m) assert(Monitor.get == m) } } - "Monitor.getOption, Monitor.setOption" should { - "maintain current monitor" in Monitor.restoring { + test("Monitor.getOption, Monitor.setOption should maintain current monitor") { + Monitor.restoring { val mon = new Monitor { def handle(cause: Throwable): Boolean = true } @@ -140,8 +130,10 @@ class MonitorTest extends WordSpec with MockitoSugar { assert(Monitor.getOption.contains(mon)) assert(mon == Monitor.get) } + } - "setOption of None restores the default" in Monitor.restoring { + test("Monitor.getOption, Monitor.setOption should setOption of None restores the default") { + Monitor.restoring { val mon = new Monitor { def handle(cause: Throwable): Boolean = true } @@ -154,10 +146,9 @@ class MonitorTest extends WordSpec with MockitoSugar { } } - "Monitor.handle" should { + test("Monitor.handle should dispatch to current monitor") { val m = spy(new MockMonitor) - - "dispatch to current monitor" in Monitor.restoring { + Monitor.restoring { when(m.handle(any[Throwable])).thenReturn(true) val exc = new Exception Monitor.set(m) @@ -166,52 +157,48 @@ class MonitorTest extends WordSpec with MockitoSugar { } } - "Monitor.restore" should { - "restore current configuration" in { - val orig = Monitor.get - Monitor.restoring { - Monitor.set(mock[Monitor]) - } - assert(Monitor.get == orig) + test("Monitor.restore should restore current configuration") { + val orig = Monitor.get + Monitor.restoring { + Monitor.set(mock[Monitor]) } + assert(Monitor.get == orig) } - "Monitor.mk" should { - class MonitorMkHelper { - class E1 extends Exception - class E2 extends E1 - class F1 extends Exception + class MonitorMkHelper { + class E1 extends Exception + class E2 extends E1 + class F1 extends Exception - var ran = false - val m = Monitor.mk { - case _: E1 => - ran = true - true - } + var ran = false + val m = Monitor.mk { + case _: E1 => + ran = true + true } + } - "handle E1" in { - val h = new MonitorMkHelper - import h._ + test("Monitor.mk should handle E1") { + val h = new MonitorMkHelper + import h._ - assert(m.handle(new E1)) - assert(ran) - } + assert(m.handle(new E1)) + assert(ran) + } - "handle E2" in { - val h = new MonitorMkHelper - import h._ + test("Monitor.mk should handle E2") { + val h = new MonitorMkHelper + import h._ - assert(m.handle(new E2)) - assert(ran) - } + assert(m.handle(new E2)) + assert(ran) + } - "not handle F1" in { - val h = new MonitorMkHelper - import h._ + test("Monitor.mk should handle F1") { + val h = new MonitorMkHelper + import h._ - assert(!m.handle(new F1)) - assert(!ran) - } + assert(!m.handle(new F1)) + assert(!ran) } } diff --git a/util-core/src/test/scala/com/twitter/util/NetUtilTest.scala b/util-core/src/test/scala/com/twitter/util/NetUtilTest.scala index 8018fbc9a5..2ea31df956 100644 --- a/util-core/src/test/scala/com/twitter/util/NetUtilTest.scala +++ b/util-core/src/test/scala/com/twitter/util/NetUtilTest.scala @@ -1,150 +1,144 @@ package com.twitter.util import java.net.InetAddress - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class NetUtilTest extends WordSpec { - "NetUtil" should { - "isIpv4Address" in { - for (i <- 0.to(255)) { - assert(NetUtil.isIpv4Address("%d.0.0.0".format(i)) == true) - assert(NetUtil.isIpv4Address("0.%d.0.0".format(i)) == true) - assert(NetUtil.isIpv4Address("0.0.%d.0".format(i)) == true) - assert(NetUtil.isIpv4Address("0.0.0.%d".format(i)) == true) - assert(NetUtil.isIpv4Address("%d.%d.%d.%d".format(i, i, i, i)) == true) - } - - assert(NetUtil.isIpv4Address("") == false) - assert(NetUtil.isIpv4Address("no") == false) - assert(NetUtil.isIpv4Address("::127.0.0.1") == false) - assert(NetUtil.isIpv4Address("-1.0.0.0") == false) - assert(NetUtil.isIpv4Address("256.0.0.0") == false) - assert(NetUtil.isIpv4Address("0.256.0.0") == false) - assert(NetUtil.isIpv4Address("0.0.256.0") == false) - assert(NetUtil.isIpv4Address("0.0.0.256") == false) - assert(NetUtil.isIpv4Address("x1.2.3.4") == false) - assert(NetUtil.isIpv4Address("1.x2.3.4") == false) - assert(NetUtil.isIpv4Address("1.2.x3.4") == false) - assert(NetUtil.isIpv4Address("1.2.3.x4") == false) - assert(NetUtil.isIpv4Address("1.2.3.4x") == false) - assert(NetUtil.isIpv4Address(" 1.2.3.4") == false) - assert(NetUtil.isIpv4Address("1.2.3.4 ") == false) - assert(NetUtil.isIpv4Address(".") == false) - assert(NetUtil.isIpv4Address("....") == false) - assert(NetUtil.isIpv4Address("1....") == false) - assert(NetUtil.isIpv4Address("1.2...") == false) - assert(NetUtil.isIpv4Address("1.2.3.") == false) - assert(NetUtil.isIpv4Address(".2.3.4") == false) +import org.scalatest.funsuite.AnyFunSuite + +class NetUtilTest extends AnyFunSuite { + test("NetUtil isIpv4Address") { + for (i <- 0.to(255)) { + assert(NetUtil.isIpv4Address("%d.0.0.0".format(i)) == true) + assert(NetUtil.isIpv4Address("0.%d.0.0".format(i)) == true) + assert(NetUtil.isIpv4Address("0.0.%d.0".format(i)) == true) + assert(NetUtil.isIpv4Address("0.0.0.%d".format(i)) == true) + assert(NetUtil.isIpv4Address("%d.%d.%d.%d".format(i, i, i, i)) == true) } - "isPrivate" in { - assert(NetUtil.isPrivateAddress(InetAddress.getByName("0.0.0.0")) == false) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("199.59.148.13")) == false) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("10.0.0.0")) == true) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("10.255.255.255")) == true) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("172.16.0.0")) == true) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("172.31.255.255")) == true) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("192.168.0.0")) == true) - assert(NetUtil.isPrivateAddress(InetAddress.getByName("192.168.255.255")) == true) - } - - "ipToInt" in { - assert(NetUtil.ipToInt("0.0.0.0") == 0) - assert(NetUtil.ipToInt("255.255.255.255") == 0xFFFFFFFF) - assert(NetUtil.ipToInt("255.255.255.0") == 0xFFFFFF00) - assert(NetUtil.ipToInt("255.0.255.0") == 0xFF00FF00) - assert(NetUtil.ipToInt("61.197.253.56") == 0x3dc5fd38) - intercept[IllegalArgumentException] { - NetUtil.ipToInt("256.0.255.0") - } - } + assert(NetUtil.isIpv4Address("") == false) + assert(NetUtil.isIpv4Address("no") == false) + assert(NetUtil.isIpv4Address("::127.0.0.1") == false) + assert(NetUtil.isIpv4Address("-1.0.0.0") == false) + assert(NetUtil.isIpv4Address("256.0.0.0") == false) + assert(NetUtil.isIpv4Address("0.256.0.0") == false) + assert(NetUtil.isIpv4Address("0.0.256.0") == false) + assert(NetUtil.isIpv4Address("0.0.0.256") == false) + assert(NetUtil.isIpv4Address("x1.2.3.4") == false) + assert(NetUtil.isIpv4Address("1.x2.3.4") == false) + assert(NetUtil.isIpv4Address("1.2.x3.4") == false) + assert(NetUtil.isIpv4Address("1.2.3.x4") == false) + assert(NetUtil.isIpv4Address("1.2.3.4x") == false) + assert(NetUtil.isIpv4Address(" 1.2.3.4") == false) + assert(NetUtil.isIpv4Address("1.2.3.4 ") == false) + assert(NetUtil.isIpv4Address(".") == false) + assert(NetUtil.isIpv4Address("....") == false) + assert(NetUtil.isIpv4Address("1....") == false) + assert(NetUtil.isIpv4Address("1.2...") == false) + assert(NetUtil.isIpv4Address("1.2.3.") == false) + assert(NetUtil.isIpv4Address(".2.3.4") == false) + } - "inetAddressToInt" in { - assert(NetUtil.inetAddressToInt(InetAddress.getByName("0.0.0.0")) == 0) - assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.255.255.255")) == 0xFFFFFFFF) - assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.255.255.0")) == 0xFFFFFF00) - assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.0.255.0")) == 0xFF00FF00) - assert(NetUtil.inetAddressToInt(InetAddress.getByName("61.197.253.56")) == 0x3dc5fd38) - intercept[IllegalArgumentException] { - NetUtil.inetAddressToInt(InetAddress.getByName("::1")) - } - } + test("NetUtil isPrivate") { + assert(NetUtil.isPrivateAddress(InetAddress.getByName("0.0.0.0")) == false) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("199.59.148.13")) == false) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("10.0.0.0")) == true) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("10.255.255.255")) == true) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("172.16.0.0")) == true) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("172.31.255.255")) == true) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("192.168.0.0")) == true) + assert(NetUtil.isPrivateAddress(InetAddress.getByName("192.168.255.255")) == true) + } - "cidrToIpBlock" in { - assert(NetUtil.cidrToIpBlock("127") == ((0x7F000000, 0xFF000000))) - assert(NetUtil.cidrToIpBlock("127.0.0") == ((0x7F000000, 0xFFFFFF00))) - assert(NetUtil.cidrToIpBlock("127.0.0.1") == ((0x7F000001, 0xFFFFFFFF))) - assert(NetUtil.cidrToIpBlock("127.0.0.1/1") == ((0x7F000001, 0x80000000))) - assert(NetUtil.cidrToIpBlock("127.0.0.1/4") == ((0x7F000001, 0xF0000000))) - assert(NetUtil.cidrToIpBlock("127.0.0.1/32") == ((0x7F000001, 0xFFFFFFFF))) - assert(NetUtil.cidrToIpBlock("127/24") == ((0x7F000000, 0xFFFFFF00))) + test("NetUtil ipToInt") { + assert(NetUtil.ipToInt("0.0.0.0") == 0) + assert(NetUtil.ipToInt("255.255.255.255") == 0xffffffff) + assert(NetUtil.ipToInt("255.255.255.0") == 0xffffff00) + assert(NetUtil.ipToInt("255.0.255.0") == 0xff00ff00) + assert(NetUtil.ipToInt("61.197.253.56") == 0x3dc5fd38) + intercept[IllegalArgumentException] { + NetUtil.ipToInt("256.0.255.0") } + } - "isInetAddressInBlock" in { - val block = NetUtil.cidrToIpBlock("192.168.0.0/16") - - assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.168.0.1"), block) == true) - assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.168.255.254"), block) == true) - assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.169.0.1"), block) == false) + test("NetUtil inetAddressToInt") { + assert(NetUtil.inetAddressToInt(InetAddress.getByName("0.0.0.0")) == 0) + assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.255.255.255")) == 0xffffffff) + assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.255.255.0")) == 0xffffff00) + assert(NetUtil.inetAddressToInt(InetAddress.getByName("255.0.255.0")) == 0xff00ff00) + assert(NetUtil.inetAddressToInt(InetAddress.getByName("61.197.253.56")) == 0x3dc5fd38) + intercept[IllegalArgumentException] { + NetUtil.inetAddressToInt(InetAddress.getByName("::1")) } + } - "isIpInBlocks" in { - val blocks = Seq( - NetUtil.cidrToIpBlock("127"), - NetUtil.cidrToIpBlock("10.1.1.0/24"), - NetUtil.cidrToIpBlock("192.168.0.0/16"), - NetUtil.cidrToIpBlock("200.1.1.1"), - NetUtil.cidrToIpBlock("200.1.1.2/32") - ) - - assert(NetUtil.isIpInBlocks("127.0.0.1", blocks) == true) - assert(NetUtil.isIpInBlocks("128.0.0.1", blocks) == false) - assert(NetUtil.isIpInBlocks("127.255.255.255", blocks) == true) - - assert(NetUtil.isIpInBlocks("10.1.1.1", blocks) == true) - assert(NetUtil.isIpInBlocks("10.1.1.255", blocks) == true) - assert(NetUtil.isIpInBlocks("10.1.0.255", blocks) == false) - assert(NetUtil.isIpInBlocks("10.1.2.0", blocks) == false) + test("NetUtil cidrToIpBlock") { + assert(NetUtil.cidrToIpBlock("127") == ((0x7f000000, 0xff000000))) + assert(NetUtil.cidrToIpBlock("127.0.0") == ((0x7f000000, 0xffffff00))) + assert(NetUtil.cidrToIpBlock("127.0.0.1") == ((0x7f000001, 0xffffffff))) + assert(NetUtil.cidrToIpBlock("127.0.0.1/1") == ((0x7f000001, 0x80000000))) + assert(NetUtil.cidrToIpBlock("127.0.0.1/4") == ((0x7f000001, 0xf0000000))) + assert(NetUtil.cidrToIpBlock("127.0.0.1/32") == ((0x7f000001, 0xffffffff))) + assert(NetUtil.cidrToIpBlock("127/24") == ((0x7f000000, 0xffffff00))) + } - assert(NetUtil.isIpInBlocks("192.168.0.1", blocks) == true) - assert(NetUtil.isIpInBlocks("192.168.255.255", blocks) == true) - assert(NetUtil.isIpInBlocks("192.167.255.255", blocks) == false) - assert(NetUtil.isIpInBlocks("192.169.0.0", blocks) == false) - assert(NetUtil.isIpInBlocks("200.168.0.0", blocks) == false) + test("NetUtil isInetAddressInBlock") { + val block = NetUtil.cidrToIpBlock("192.168.0.0/16") - assert(NetUtil.isIpInBlocks("200.1.1.1", blocks) == true) - assert(NetUtil.isIpInBlocks("200.1.3.1", blocks) == false) - assert(NetUtil.isIpInBlocks("200.1.1.2", blocks) == true) - assert(NetUtil.isIpInBlocks("200.1.3.2", blocks) == false) + assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.168.0.1"), block) == true) + assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.168.255.254"), block) == true) + assert(NetUtil.isInetAddressInBlock(InetAddress.getByName("192.169.0.1"), block) == false) + } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("no", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("::127.0.0.1", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("-1.0.0.0", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("256.0.0.0", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("0.256.0.0", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("0.0.256.0", blocks) - } - intercept[IllegalArgumentException] { - NetUtil.isIpInBlocks("0.0.0.256", blocks) - } + test("NetUtil isIpInBlocks") { + val blocks = Seq( + NetUtil.cidrToIpBlock("127"), + NetUtil.cidrToIpBlock("10.1.1.0/24"), + NetUtil.cidrToIpBlock("192.168.0.0/16"), + NetUtil.cidrToIpBlock("200.1.1.1"), + NetUtil.cidrToIpBlock("200.1.1.2/32") + ) + + assert(NetUtil.isIpInBlocks("127.0.0.1", blocks) == true) + assert(NetUtil.isIpInBlocks("128.0.0.1", blocks) == false) + assert(NetUtil.isIpInBlocks("127.255.255.255", blocks) == true) + + assert(NetUtil.isIpInBlocks("10.1.1.1", blocks) == true) + assert(NetUtil.isIpInBlocks("10.1.1.255", blocks) == true) + assert(NetUtil.isIpInBlocks("10.1.0.255", blocks) == false) + assert(NetUtil.isIpInBlocks("10.1.2.0", blocks) == false) + + assert(NetUtil.isIpInBlocks("192.168.0.1", blocks) == true) + assert(NetUtil.isIpInBlocks("192.168.255.255", blocks) == true) + assert(NetUtil.isIpInBlocks("192.167.255.255", blocks) == false) + assert(NetUtil.isIpInBlocks("192.169.0.0", blocks) == false) + assert(NetUtil.isIpInBlocks("200.168.0.0", blocks) == false) + + assert(NetUtil.isIpInBlocks("200.1.1.1", blocks) == true) + assert(NetUtil.isIpInBlocks("200.1.3.1", blocks) == false) + assert(NetUtil.isIpInBlocks("200.1.1.2", blocks) == true) + assert(NetUtil.isIpInBlocks("200.1.3.2", blocks) == false) + + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("no", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("::127.0.0.1", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("-1.0.0.0", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("256.0.0.0", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("0.256.0.0", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("0.0.256.0", blocks) + } + intercept[IllegalArgumentException] { + NetUtil.isIpInBlocks("0.0.0.256", blocks) } } } diff --git a/util-core/src/test/scala/com/twitter/util/PoolTest.scala b/util-core/src/test/scala/com/twitter/util/PoolTest.scala index fd11649556..71cb107617 100644 --- a/util-core/src/test/scala/com/twitter/util/PoolTest.scala +++ b/util-core/src/test/scala/com/twitter/util/PoolTest.scala @@ -1,75 +1,68 @@ package com.twitter.util +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite import scala.collection.mutable -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +class PoolTest extends AnyFunSuite { -import com.twitter.conversions.time._ + private[this] def await[T](f: Awaitable[T]): T = + Await.result(f, 2.seconds) -@RunWith(classOf[JUnitRunner]) -class PoolTest extends WordSpec { - "SimplePool" should { - "with a simple queue of items" should { - "it reseves items in FIFO order" in { - val queue = new mutable.Queue[Int] ++ List(1, 2, 3) - val pool = new SimplePool(queue) - assert(Await.result(pool.reserve()) == 1) - assert(Await.result(pool.reserve()) == 2) - pool.release(2) - assert(Await.result(pool.reserve()) == 3) - assert(Await.result(pool.reserve()) == 2) - pool.release(1) - pool.release(2) - pool.release(3) - } + class PoolSpecHelper { + var count = 0 + val pool = new FactoryPool[Int](4) { + def makeItem() = { count += 1; Future(count) } + def isHealthy(i: Int) = i % 2 == 0 } + } - "with an object factory and a health check" should { - class PoolSpecHelper { - var count = 0 - val pool = new FactoryPool[Int](4) { - def makeItem() = { count += 1; Future(count) } - def isHealthy(i: Int) = i % 2 == 0 - } - } + test("SimplePool with a simple queue of items should reserve items in FIFO order") { + val queue = new mutable.Queue[Int] ++ List(1, 2, 3) + val pool = new SimplePool(queue) + assert(await(pool.reserve()) == 1) + assert(await(pool.reserve()) == 2) + pool.release(2) + assert(await(pool.reserve()) == 3) + assert(await(pool.reserve()) == 2) + pool.release(1) + pool.release(2) + pool.release(3) + } - "reserve & release" in { - val h = new PoolSpecHelper - import h._ + test("SimplePool with an object factory and a health check should reserve & release") { + val h = new PoolSpecHelper + import h._ - assert(Await.result(pool.reserve()) == 2) - assert(Await.result(pool.reserve()) == 4) - assert(Await.result(pool.reserve()) == 6) - assert(Await.result(pool.reserve()) == 8) - val promise = pool.reserve() - intercept[TimeoutException] { - Await.result(promise, 1.millisecond) - } - pool.release(8) - pool.release(6) - assert(Await.result(promise) == 8) - assert(Await.result(pool.reserve()) == 6) - intercept[TimeoutException] { - Await.result(pool.reserve(), 1.millisecond) - } - } + assert(await(pool.reserve()) == 2) + assert(await(pool.reserve()) == 4) + assert(await(pool.reserve()) == 6) + assert(await(pool.reserve()) == 8) + val promise = pool.reserve() + intercept[TimeoutException] { + Await.result(promise, 1.millisecond) + } + pool.release(8) + pool.release(6) + assert(await(promise) == 8) + assert(await(pool.reserve()) == 6) + intercept[TimeoutException] { + Await.result(pool.reserve(), 1.millisecond) + } + } - "reserve & dispose" in { - val h = new PoolSpecHelper - import h._ + test("SimplePool with an object factory and a health check should reserve & dispose") { + val h = new PoolSpecHelper + import h._ - assert(Await.result(pool.reserve()) == 2) - assert(Await.result(pool.reserve()) == 4) - assert(Await.result(pool.reserve()) == 6) - assert(Await.result(pool.reserve()) == 8) - intercept[TimeoutException] { - Await.result(pool.reserve(), 1.millisecond) - } - pool.dispose(2) - assert(Await.result(pool.reserve(), 1.millisecond) == 10) - } + assert(await(pool.reserve()) == 2) + assert(await(pool.reserve()) == 4) + assert(await(pool.reserve()) == 6) + assert(await(pool.reserve()) == 8) + intercept[TimeoutException] { + Await.result(pool.reserve(), 1.millisecond) } + pool.dispose(2) + assert(Await.result(pool.reserve(), 1.millisecond) == 10) } } diff --git a/util-core/src/test/scala/com/twitter/util/PromiseTest.scala b/util-core/src/test/scala/com/twitter/util/PromiseTest.scala index 9f28e5a554..7eb2e68b43 100644 --- a/util-core/src/test/scala/com/twitter/util/PromiseTest.scala +++ b/util-core/src/test/scala/com/twitter/util/PromiseTest.scala @@ -1,12 +1,80 @@ package com.twitter.util -import com.twitter.conversions.time._ -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.concurrent.Scheduler +import com.twitter.concurrent.LocalScheduler +import com.twitter.conversions.DurationOps._ +import java.util.concurrent.atomic.AtomicBoolean +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class PromiseTest extends FunSuite { +class PromiseTest extends AnyFunSuite { + + final class LinkablePromise[A] extends Promise[A] { + def link1(target: Promise[A]): Unit = link(target) + } + + class TrackingScheduler extends LocalScheduler(false) { + @volatile var submitted: Int = 0 + override def submit(r: Runnable): Unit = { + submitted = submitted + 1 + super.submit(r) + } + } + + class MockTracker(var updates: Long = 0) extends ResourceTracker { + override def addCpuTime(time: Long): Unit = { + updates = updates + 1 + } + } + + test("ResourceTracker when set, updates per completed continuation") { + if (ResourceTracker.threadCpuTimeSupported) { + val mockTracker = new MockTracker() + val p1 = new Promise[Unit]() + val p = ResourceTracker.let(mockTracker) { + // 2 continuations + p1.map(_ => ()).respond(_ => ()) + + // 1 continuation + // + // **Note**: Because [[ConstFuture]] implements `transformTry` synchronously, + // the cpuTime is not updated since on `Future.Unit#map`, instead accounted for + // by the saved context of the parent continuation, `flatMap`. + p1.flatMap(_ => Future.Unit.map(_ => ())) + } + p1.setDone() + assert(mockTracker.updates == 3) + } + } + + test("ResourceTracker updates in the presence of exceptions") { + if (ResourceTracker.threadCpuTimeSupported) { + val mockTracker = new MockTracker() + val p = ResourceTracker.let(mockTracker) { + Future.Unit.transform(_ => throw new Exception("boom again")) + } + assert(mockTracker.updates == 1) + } + } + + test("ResourceTracker handles nested promise") { + if (ResourceTracker.threadCpuTimeSupported) { + val mockTracker = new MockTracker() + val p1 = new Promise[Unit]() + val p2 = new Promise[Unit]() + + val p = ResourceTracker.let(mockTracker) { + p1.flatMap { _ => + p2.map(_ => ()) + } + } + + p1.setDone() + assert(mockTracker.updates == 1) + + p2.setDone(); Await.result(p) + assert(mockTracker.updates == 2) + } + } test("Promise.detach should not detach other attached promises") { val p = new Promise[Unit] @@ -97,6 +165,58 @@ class PromiseTest extends FunSuite { assert(ex.getMessage.contains(value)) } + test("interrupted.link(done) runs interrupted waitqueue") { + val q = new LinkablePromise[String] // interrupted with a waitqueue + val run = new AtomicBoolean(false) + q.respond(_ => run.set(true)) + q.raise(new Exception("boom")) + + val p = new Promise[String](Return("foo")) // satisfied + + q.link1(p) + assert(run.get()) + } + + test("transforming(never).link(done) runs transforming waitqueue") { + val q = new LinkablePromise[String] // interrupted with a waitqueue + val deadend = Future.never + val run = new AtomicBoolean(false) + q.respond(_ => run.set(true)) + q.forwardInterruptsTo(deadend) + + val p = new Promise[String](Return("foo")) // satisfied + + q.link1(p) + assert(run.get()) + } + + test("transforming(done).link(done) runs transforming waitqueue") { + val q = new LinkablePromise[String] // interrupted with a waitqueue + val deadend = new Promise[Unit] + val run = new AtomicBoolean(false) + q.respond(_ => run.set(true)) + q.forwardInterruptsTo(deadend) + + val p = new Promise[String](Return("foo")) // satisfied + deadend.setDone() + + q.link1(p) + assert(run.get()) + } + + test("interruptible.link(done) runs interruptible waitqueue") { + val q = new LinkablePromise[String] // interrupted with a waitqueue + val deadend = Future.Done + val run = new AtomicBoolean(false) + q.respond(_ => run.set(true)) + q.setInterruptHandler { case _ => () } + + val p = new Promise[String](Return("foo")) // satisfied + + q.link1(p) + assert(run.get()) + } + test("Updating a Promise more than once should fail") { val p = new Promise[Int]() val first = Return(1) @@ -128,8 +248,8 @@ class PromiseTest extends FunSuite { } class HandledMonitor extends Monitor { - var handled = null: Throwable - def handle(exc: Throwable) = { + var handled: Throwable = null: Throwable + def handle(exc: Throwable): Boolean = { handled = exc true } @@ -198,9 +318,142 @@ class PromiseTest extends FunSuite { test("ready") { // this verifies that we don't break the behavior regarding how // the `Scheduler.flush()` is a necessary in `Promise.ready` - Future.Done.map { _ => - Await.result(Future.Done.map(Predef.identity), 5.seconds) + Future.Done.map { _ => Await.result(Future.Done.map(Predef.identity), 5.seconds) } + } + + test("maps use local state at the time of callback creation") { + val local = new Local[String] + local.set(Some("main thread")) + val p = new Promise[String]() + val f = p.map { _ => local().getOrElse("") } + + val thread = new Thread(new Runnable { + def run(): Unit = { + local.set(Some("worker thread")) + p.setValue("go") + } + }) + thread.start() + + assert("main thread" == Await.result(f, 3.seconds)) + } + + test("interrupts use local state at the time of callback creation") { + val local = new Local[String] + implicit val timer: Timer = new JavaTimer(true) + + local.set(Some("main thread")) + val p = new Promise[String] + + p.setInterruptHandler({ + case _ => + val localVal = local().getOrElse("") + p.updateIfEmpty(Return(localVal)) + }) + val f = Future.sleep(1.second).before(p) + + val thread = new Thread(new Runnable { + def run(): Unit = { + local.set(Some("worker thread")) + p.raise(new Exception) + } + }) + thread.start() + + assert("main thread" == Await.result(f, 3.seconds)) + } + + test("interrupt then map uses local state at the time of callback creation") { + val local = new Local[String] + implicit val timer: Timer = new JavaTimer(true) + + local.set(Some("main thread")) + val p = new Promise[String]() + p.setInterruptHandler({ + case _ => + val localVal = local().getOrElse("") + p.updateIfEmpty(Return(localVal)) + }) + val fRes = p.transform { _ => + val localVal = local().getOrElse("") + Future.value(localVal) + } + val f = Future.sleep(1.second).before(fRes) + + val thread = new Thread(new Runnable { + def run(): Unit = { + local.set(Some("worker thread")) + p.raise(new Exception) + } + }) + thread.start() + + assert("main thread" == Await.result(f, 3.seconds)) + } + + test("locals in Interruptible after detatch") { + val local = new Local[String] + implicit val timer: Timer = new JavaTimer(true) + + local.set(Some("main thread")) + val p = new Promise[String] + + val p2 = Promise.attached(p) // p2 is the detachable promise + + p2.setInterruptHandler({ + case _ => + val localVal = local().getOrElse("") + p2.updateIfEmpty(Return(localVal)) + }) + + p.setInterruptHandler({ + case _ => + val localVal = local().getOrElse("") + p.updateIfEmpty(Return(localVal)) + }) + + val thread = new Thread(new Runnable { + def run(): Unit = { + local.set(Some("worker thread")) + + // We need to call detach explicitly + assert(p2.detach()) + + p2.raise(new Exception) + p.raise(new Exception) + } + }) + thread.start() + + assert("main thread" == Await.result(p2, 3.seconds)) + assert("main thread" == Await.result(p, 3.seconds)) + } + + test("promise with no continuations is not submitted to the scheduler") { + val old = Scheduler() + try { + val sched = new TrackingScheduler() + Scheduler.setUnsafe(sched) + val p = new Promise[Unit]() + p.setDone() + assert(sched.submitted == 0) + } finally { + Scheduler.setUnsafe(old) } } + test("registered continuation on satisfied promise is submitted to the scheduler") { + val old = Scheduler() + try { + val sched = new TrackingScheduler() + Scheduler.setUnsafe(sched) + val p = new Promise[Unit]() + p.setDone() + + p.ensure(()) + assert(sched.submitted == 1) + } finally { + Scheduler.setUnsafe(old) + } + } } diff --git a/util-core/src/test/scala/com/twitter/util/SchedulerWorkQueueFiberTest.scala b/util-core/src/test/scala/com/twitter/util/SchedulerWorkQueueFiberTest.scala new file mode 100644 index 0000000000..e56a49cc93 --- /dev/null +++ b/util-core/src/test/scala/com/twitter/util/SchedulerWorkQueueFiberTest.scala @@ -0,0 +1,60 @@ +package com.twitter.util + +import com.twitter.concurrent.LocalScheduler +import java.util.concurrent.CountDownLatch +import org.scalatest.concurrent.Eventually.eventually +import org.scalatest.concurrent.PatienceConfiguration.Timeout +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.time.Seconds +import org.scalatest.time.Span + +class SchedulerWorkQueueFiberTest extends AnyFunSuite { + + private class NoOpMetrics extends WorkQueueFiber.FiberMetrics { + def flushIncrement(): Unit = () + def fiberCreated(): Unit = () + def threadLocalSubmitIncrement(): Unit = () + def taskSubmissionIncrement(): Unit = () + def schedulerSubmissionIncrement(): Unit = () + } + + test("fiber flushes the underlying LocalScheduler") { + val flushLatch = new CountDownLatch(1) + val localScheduler = new LocalScheduler { + override def flush(): Unit = { + super.flush() + flushLatch.countDown() + } + } + val localSchedulerFiber = new SchedulerWorkQueueFiber(localScheduler, new NoOpMetrics) + + var promiseSatisfied = false + + val runnable = new Runnable { + def run(): Unit = { + val p = Promise[Boolean]() + val f2 = new SchedulerWorkQueueFiber(localScheduler, new NoOpMetrics) + // Satisfy the promise on another fiber, whose work will get scheduled on the same thread + f2.submitTask(() => p.setValue(true)) + + // This will trigger a flush of the current fiber and the local scheduler + Await.result(p) + promiseSatisfied = true + } + } + + val t = new Thread({ () => + Fiber.let(localSchedulerFiber) { + localSchedulerFiber.submitTask(() => runnable.run()) + } + }) + t.start() + + // Ensure flush completed + flushLatch.await() + + // The promise should be satisfied even though the second fiber was scheduled on the + // same thread as the Await.result + eventually(Timeout(Span(1, Seconds))) { assert(promiseSatisfied) } + } +} diff --git a/util-core/src/test/scala/com/twitter/util/SlowProbeProxyTimerTest.scala b/util-core/src/test/scala/com/twitter/util/SlowProbeProxyTimerTest.scala index 45e8f92f40..770c3f5bc8 100644 --- a/util-core/src/test/scala/com/twitter/util/SlowProbeProxyTimerTest.scala +++ b/util-core/src/test/scala/com/twitter/util/SlowProbeProxyTimerTest.scala @@ -1,10 +1,10 @@ package com.twitter.util -import com.twitter.conversions.time._ -import org.scalatest.FunSuite +import com.twitter.conversions.DurationOps._ import scala.collection.mutable +import org.scalatest.funsuite.AnyFunSuite -class SlowProbeProxyTimerTest extends FunSuite { +class SlowProbeProxyTimerTest extends AnyFunSuite { private type Task = () => Unit private val NullTask = new TimerTask { def cancel(): Unit = () } diff --git a/util-core/src/test/scala/com/twitter/util/StateMachineTest.scala b/util-core/src/test/scala/com/twitter/util/StateMachineTest.scala index a93a416c8e..cbb541d68e 100644 --- a/util-core/src/test/scala/com/twitter/util/StateMachineTest.scala +++ b/util-core/src/test/scala/com/twitter/util/StateMachineTest.scala @@ -1,43 +1,38 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class StateMachineTest extends WordSpec { - "StateMachine" should { - class StateMachineHelper { - class Machine extends StateMachine { - case class State1() extends State - case class State2() extends State - state = State1() - - def command1(): Unit = { - transition("command1") { - case State1() => state = State2() - } +import org.scalatest.funsuite.AnyFunSuite + +class StateMachineTest extends AnyFunSuite { + class StateMachineHelper { + class Machine extends StateMachine { + case class State1() extends State + case class State2() extends State + state = State1() + + def command1(): Unit = { + transition("command1") { + case State1() => state = State2() } } - - val stateMachine = new Machine() } - "allows transitions that are permitted" in { - val h = new StateMachineHelper - import h._ + val stateMachine = new Machine() + } - stateMachine.command1() - } + test("StateMachine allows transitions that are permitted") { + val h = new StateMachineHelper + import h._ - "throws exceptions when a transition is not permitted" in { - val h = new StateMachineHelper - import h._ + stateMachine.command1() + } + + test("StateMachine throws exceptions when a transition is not permitted") { + val h = new StateMachineHelper + import h._ + stateMachine.command1() + intercept[StateMachine.InvalidStateTransition] { stateMachine.command1() - intercept[StateMachine.InvalidStateTransition] { - stateMachine.command1() - } } } } diff --git a/util-core/src/test/scala/com/twitter/util/StorageUnitTest.scala b/util-core/src/test/scala/com/twitter/util/StorageUnitTest.scala index 3dbac51527..a057a6ae1d 100644 --- a/util-core/src/test/scala/com/twitter/util/StorageUnitTest.scala +++ b/util-core/src/test/scala/com/twitter/util/StorageUnitTest.scala @@ -1,27 +1,32 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.conversions.StorageUnitOps._ +import org.scalatest.funsuite.AnyFunSuite -import com.twitter.conversions.storage._ - -@RunWith(classOf[JUnitRunner]) -class StorageUnitTest extends FunSuite { +class StorageUnitTest extends AnyFunSuite { test("StorageUnit: should convert whole numbers into storage units (back and forth)") { assert(1.byte.inBytes == 1) assert(1.kilobyte.inBytes == 1024) assert(1.megabyte.inMegabytes == 1.0) assert(1.gigabyte.inMegabytes == 1024.0) assert(1.gigabyte.inKilobytes == 1024.0 * 1024.0) + assert(1.terabytes.inTerabytes == 1) + assert(1.terabytes.inGigabytes == 1024.0) + assert(1.terabytes.inMegabytes == 1024.0 * 1024.0) + assert(1.petabytes.inPetabytes == 1) + assert(1.petabytes.inTerabytes == 1024.0) + assert(1.petabytes.inGigabytes == 1024.0 * 1024.0) + assert((1.petabytes * 1024.0).inExabytes == 1) + assert((1.petabytes * 1024.0).inPetabytes == 1024.0) + assert((1.petabytes * 1024.0).inTerabytes == 1024.0 * 1024.0) } test("StorageUnit: should confer an essential humanity") { - assert(900.bytes.toHuman == "900 B") - assert(1.kilobyte.toHuman == "1024 B") - assert(2.kilobytes.toHuman == "2.0 KiB") - assert(Int.MaxValue.bytes.toHuman == "2.0 GiB") - assert(Long.MaxValue.bytes.toHuman == "8.0 EiB") + assert(900.bytes.toHuman() == "900 B") + assert(1.kilobyte.toHuman() == "1024 B") + assert(2.kilobytes.toHuman() == "2.0 KiB") + assert(Int.MaxValue.bytes.toHuman() == "2.0 GiB") + assert(Long.MaxValue.bytes.toHuman() == "8.0 EiB") } test("StorageUnit: should handle Long value") { @@ -36,6 +41,7 @@ class StorageUnitTest extends FunSuite { assert(StorageUnit.parse("3.terabytes") == 3.terabytes) assert(StorageUnit.parse("9.petabytes") == 9.petabytes) assert(StorageUnit.parse("-3.megabytes") == -3.megabytes) + assert(StorageUnit.parse("3.exabytes") == 3.petabytes * 1024.0) } test("StorageUnit: should reject soulless robots") { @@ -45,7 +51,7 @@ class StorageUnitTest extends FunSuite { test("StorageUnit: should deal with negative values") { assert(-123.bytes.inBytes == -123) - assert(-2.kilobytes.toHuman == "-2.0 KiB") + assert(-2.kilobytes.toHuman() == "-2.0 KiB") } test("StorageUnit: should min properly") { diff --git a/util-core/src/test/scala/com/twitter/util/StringConversionsTest.scala b/util-core/src/test/scala/com/twitter/util/StringConversionsTest.scala deleted file mode 100644 index b2e0baabeb..0000000000 --- a/util-core/src/test/scala/com/twitter/util/StringConversionsTest.scala +++ /dev/null @@ -1,60 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.util - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -import com.twitter.conversions.string._ - -@RunWith(classOf[JUnitRunner]) -class StringConversionsTest extends WordSpec { - "string" should { - "quoteC" in { - assert("nothing".quoteC == "nothing") - assert( - "name\tvalue\t\u20acb\u00fcllet?\u20ac".quoteC == "name\\tvalue\\t\\u20acb\\xfcllet?\\u20ac" - ) - assert("she said \"hello\"".quoteC == "she said \\\"hello\\\"") - assert("\\backslash".quoteC == "\\\\backslash") - } - - "unquoteC" in { - assert("nothing".unquoteC == "nothing") - assert( - "name\\tvalue\\t\\u20acb\\xfcllet?\\u20ac".unquoteC == "name\tvalue\t\u20acb\u00fcllet?\u20ac" - ) - assert("she said \\\"hello\\\"".unquoteC == "she said \"hello\"") - assert("\\\\backslash".unquoteC == "\\backslash") - assert("real\\$dollar".unquoteC == "real\\$dollar") - assert("silly\\/quote".unquoteC == "silly/quote") - } - - "hexlify" in { - assert("hello".getBytes.slice(1, 4).hexlify == "656c6c") - assert("hello".getBytes.hexlify == "68656c6c6f") - } - - "unhexlify" in { - assert("656c6c".unhexlify.toList == "hello".getBytes.slice(1, 4).toList) - assert("68656c6c6f".unhexlify.toList == "hello".getBytes.toList) - "5".unhexlify - assert("5".unhexlify.hexlify.toInt == 5) - } - } -} diff --git a/util-core/src/test/scala/com/twitter/util/ThrowablesTest.scala b/util-core/src/test/scala/com/twitter/util/ThrowablesTest.scala index 626e404750..0661094553 100644 --- a/util-core/src/test/scala/com/twitter/util/ThrowablesTest.scala +++ b/util-core/src/test/scala/com/twitter/util/ThrowablesTest.scala @@ -1,11 +1,8 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ThrowablesTest extends FunSuite { +class ThrowablesTest extends AnyFunSuite { test("Throwables.mkString: flatten non-nested exception") { val t = new Throwable assert(Throwables.mkString(t) == Seq(t.getClass.getName)) diff --git a/util-core/src/test/scala/com/twitter/util/TimeTest.scala b/util-core/src/test/scala/com/twitter/util/TimeTest.scala index 60e17819ad..b757e0d795 100644 --- a/util-core/src/test/scala/com/twitter/util/TimeTest.scala +++ b/util-core/src/test/scala/com/twitter/util/TimeTest.scala @@ -1,685 +1,825 @@ package com.twitter.util -import java.io.{ByteArrayInputStream, ByteArrayOutputStream, ObjectInputStream, ObjectOutputStream} -import java.util.{Locale, TimeZone} +import com.twitter.conversions.DurationOps._ +import java.io.ByteArrayInputStream +import java.io.ByteArrayOutputStream +import java.io.ObjectInputStream +import java.io.ObjectOutputStream +import java.time.Instant +import java.time.OffsetDateTime +import java.time.ZoneId +import java.time.ZoneOffset +import java.time.ZonedDateTime +import java.util.Locale +import java.util.TimeZone import java.util.concurrent.TimeUnit +import org.scalatest.concurrent.Eventually +import org.scalatest.concurrent.IntegrationPatience +import org.scalatest.funsuite.AnyFunSuite +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.concurrent.{Eventually, IntegrationPatience} -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks - -import com.twitter.util.TimeConversions._ - -trait TimeLikeSpec[T <: TimeLike[T]] extends WordSpec with GeneratorDrivenPropertyChecks { +trait TimeLikeSpec[T <: TimeLike[T]] extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { val ops: TimeLikeOps[T] import ops._ - "Top, Bottom, Undefined, Nanoseconds(_), Finite(_)" should { - val easyVs = Seq(Zero, Top, Bottom, Undefined, fromNanoseconds(1), fromNanoseconds(-1)) - val vs = easyVs ++ Seq(fromNanoseconds(Long.MaxValue - 1), fromNanoseconds(Long.MinValue + 1)) - - "behave like boxed doubles" in { - assert((Top compare Undefined) < 0) - assert((Bottom compare Top) < 0) - assert((Undefined compare Undefined) == 0) - assert((Top compare Top) == 0) - assert((Bottom compare Bottom) == 0) - - assert(Top + Duration.Top == Top) - assert(Bottom - Duration.Bottom == Undefined) - assert(Top - Duration.Top == Undefined) - assert(Bottom + Duration.Bottom == Bottom) - } + val easyVs = Seq(Zero, Top, Bottom, Undefined, fromNanoseconds(1), fromNanoseconds(-1)) + val vs = easyVs ++ Seq(fromNanoseconds(Long.MaxValue - 1), fromNanoseconds(Long.MinValue + 1)) - "complementary diff" in { - // Note that this doesn't always hold because of two's - // complement arithmetic. - for (a <- easyVs; b <- easyVs) - assert((a diff b) == -(b diff a)) + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should behave like boxed doubles") { + assert((Top compare Undefined) < 0) + assert((Bottom compare Top) < 0) + assert((Undefined compare Undefined) == 0) + assert((Top compare Top) == 0) + assert((Bottom compare Bottom) == 0) - } + assert(Top + Duration.Top == Top) + assert(Bottom - Duration.Bottom == Undefined) + assert(Top - Duration.Top == Undefined) + assert(Bottom + Duration.Bottom == Bottom) + } - "complementary compare" in { - for (a <- vs; b <- vs) { - val x = a compare b - val y = b compare a - assert(((x == 0 && y == 0) || (x < 0 != y < 0)) == true) - } - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should complementary diff") { + // Note that this doesn't always hold because of two's + // complement arithmetic. + for (a <- easyVs; b <- easyVs) + assert((a diff b) == -(b diff a)) - "commutative max" in { - for (a <- vs; b <- vs) - assert((a max b) == (b max a)) - } + } - "commutative min" in { - for (a <- vs; b <- vs) - assert((a min b) == (b min a)) + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should complementary compare") { + for (a <- vs; b <- vs) { + val x = a compare b + val y = b compare a + assert(((x == 0 && y == 0) || (x < 0 != y < 0)) == true) } + } - "handle underflows" in { - assert(fromNanoseconds(Long.MinValue) - 1.nanosecond == Bottom) - assert(fromMicroseconds(Long.MinValue) - 1.nanosecond == Bottom) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should commutative max") { + for (a <- vs; b <- vs) + assert((a max b) == (b max a)) + } - "handle overflows" in { - assert(fromNanoseconds(Long.MaxValue) + 1.nanosecond == Top) - assert(fromMicroseconds(Long.MaxValue) + 1.nanosecond == Top) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should commutative min") { + for (a <- vs; b <- vs) + assert((a min b) == (b min a)) + } - "Nanoseconds(_) extracts only finite values, in nanoseconds" in { - for (t <- Seq(Top, Bottom, Undefined)) - assert(t match { - case Nanoseconds(_) => false - case _ => true - }) - - for (ns <- Seq(Long.MinValue, -1, 0, 1, Long.MaxValue); t = fromNanoseconds(ns)) - assert(t match { - case Nanoseconds(`ns`) => true - case _ => false - }) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should handle underflows") { + assert(fromNanoseconds(Long.MinValue) - 1.nanosecond == Bottom) + assert(fromMicroseconds(Long.MinValue) - 1.nanosecond == Bottom) + } - "Finite(_) extracts only finite values" in { - for (t <- Seq(Top, Bottom, Undefined)) - assert(t match { - case Finite(_) => false - case _ => true - }) - - for (ns <- Seq(Long.MinValue, -1, 0, 1, Long.MaxValue); t = fromNanoseconds(ns)) - assert(t match { - case Finite(`t`) => true - case _ => false - }) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should handle overflows") { + assert(fromNanoseconds(Long.MaxValue) + 1.nanosecond == Top) + assert(fromMicroseconds(Long.MaxValue) + 1.nanosecond == Top) + } - "roundtrip through serialization" in { - for (v <- vs) { - val bytes = new ByteArrayOutputStream - val out = new ObjectOutputStream(bytes) - out.writeObject(v) - val in = new ObjectInputStream(new ByteArrayInputStream(bytes.toByteArray)) - assert(in.readObject() == v) - } - } + test( + "Top, Bottom, Undefined, Nanoseconds(_), Finite(_): Nanoseconds(_) extracts only finite values, in nanoseconds") { + for (t <- Seq(Top, Bottom, Undefined)) + assert(t match { + case Nanoseconds(_) => false + case _ => true + }) + + for (ns <- Seq(Long.MinValue, -1, 0, 1, Long.MaxValue); t = fromNanoseconds(ns)) + assert(t match { + case Nanoseconds(`ns`) => true + case _ => false + }) + } - "has correct isZero behaviour" in { - for (t <- Seq(Top, Bottom, Undefined, fromNanoseconds(1L))) { - assert(t.isZero == false) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_): Finite(_) extracts only finite values") { + for (t <- Seq(Top, Bottom, Undefined)) + assert(t match { + case Finite(_) => false + case _ => true + }) + + for (ns <- Seq(Long.MinValue, -1, 0, 1, Long.MaxValue); t = fromNanoseconds(ns)) + assert(t match { + case Finite(`t`) => true + case _ => false + }) + } - for (z <- Seq( - Zero, - fromNanoseconds(0), - fromMicroseconds(0), - fromFractionalSeconds(0), - fromMilliseconds(0), - fromSeconds(0) - )) { - assert(z.isZero == true) - } + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) should roundtrip through serialization") { + for (v <- vs) { + val bytes = new ByteArrayOutputStream + val out = new ObjectOutputStream(bytes) + out.writeObject(v) + val in = new ObjectInputStream(new ByteArrayInputStream(bytes.toByteArray)) + assert(in.readObject() == v) } } - "Top" should { - "be impermeable to finite arithmetic" in { - assert(Top - 0.seconds == Top) - assert(Top - 100.seconds == Top) - assert(Top - Duration.fromNanoseconds(Long.MaxValue) == Top) + test("Top, Bottom, Undefined, Nanoseconds(_), Finite(_) has correct isZero behaviour") { + for (t <- Seq(Top, Bottom, Undefined, fromNanoseconds(1L))) { + assert(t.isZero == false) } - "become undefined when subtracted from itself, or added to bottom" in { - assert(Top - Duration.Top == Undefined) - assert(Top + Duration.Bottom == Undefined) + for (z <- Seq( + Zero, + fromNanoseconds(0), + fromMicroseconds(0), + fromFractionalSeconds(0), + fromMilliseconds(0), + fromSeconds(0), + fromMinutes(0) + )) { + assert(z.isZero == true) } + } - "not be equal to the maximum value" in { - assert(fromNanoseconds(Long.MaxValue) != Top) - } + test("Top should be impermeable to finite arithmetic") { + assert(Top - 0.seconds == Top) + assert(Top - 100.seconds == Top) + assert(Top - Duration.fromNanoseconds(Long.MaxValue) == Top) + } - "always be max" in { - assert((Top max fromSeconds(1)) == Top) - assert((Top max fromFractionalSeconds(1.0)) == Top) - assert((Top max fromNanoseconds(Long.MaxValue)) == Top) - assert((Top max Bottom) == Top) - } + test("Top should become undefined when subtracted from itself, or added to bottom") { + assert(Top - Duration.Top == Undefined) + assert(Top + Duration.Bottom == Undefined) + } - "greater than everything else" in { - assert(fromSeconds(0) < Top) - assert(fromFractionalSeconds(Double.MaxValue) < Top) - assert(fromNanoseconds(Long.MaxValue) < Top) - } + test("Top should not be equal to the maximum value") { + assert(fromNanoseconds(Long.MaxValue) != Top) + } - "equal to itself" in { - assert(Top == Top) - } + test("Top should always be max") { + assert((Top max fromSeconds(1)) == Top) + assert((Top max fromFractionalSeconds(1.0)) == Top) + assert((Top max fromNanoseconds(Long.MaxValue)) == Top) + assert((Top max Bottom) == Top) + } - "more or less equals only to itself" in { - assert(Top.moreOrLessEquals(Top, Duration.Top) == true) - assert(Top.moreOrLessEquals(Top, Duration.Zero) == true) - assert(Top.moreOrLessEquals(Bottom, Duration.Top) == true) - assert(Top.moreOrLessEquals(Bottom, Duration.Zero) == false) - assert(Top.moreOrLessEquals(fromSeconds(0), Duration.Top) == true) - assert(Top.moreOrLessEquals(fromSeconds(0), Duration.Bottom) == false) - } + test("Top should greater than everything else") { + assert(fromSeconds(0) < Top) + assert(fromFractionalSeconds(Double.MaxValue) < Top) + assert(fromNanoseconds(Long.MaxValue) < Top) + } - "Undefined diff to Top" in { - assert((Top diff Top) == Duration.Undefined) - } + test("Top should equal to itself") { + assert(Top == Top) } - "Bottom" should { - "be impermeable to finite arithmetic" in { - assert(Bottom + 0.seconds == Bottom) - assert(Bottom + 100.seconds == Bottom) - assert(Bottom + Duration.fromNanoseconds(Long.MaxValue) == Bottom) - } + test("Top should more or less equals only to itself") { + assert(Top.moreOrLessEquals(Top, Duration.Top) == true) + assert(Top.moreOrLessEquals(Top, Duration.Zero) == true) + assert(Top.moreOrLessEquals(Bottom, Duration.Top) == true) + assert(Top.moreOrLessEquals(Bottom, Duration.Zero) == false) + assert(Top.moreOrLessEquals(fromSeconds(0), Duration.Top) == true) + assert(Top.moreOrLessEquals(fromSeconds(0), Duration.Bottom) == false) + } - "become undefined when added with Top or subtracted by bottom" in { - assert(Bottom + Duration.Top == Undefined) - assert(Bottom - Duration.Bottom == Undefined) - } + test("Top should Undefined diff to Top") { + assert((Top diff Top) == Duration.Undefined) + } - "always be min" in { - assert((Bottom min Top) == Bottom) - assert((Bottom min fromNanoseconds(0)) == Bottom) - } + test("Bottom should be impermeable to finite arithmetic") { + assert(Bottom + 0.seconds == Bottom) + assert(Bottom + 100.seconds == Bottom) + assert(Bottom + Duration.fromNanoseconds(Long.MaxValue) == Bottom) + } - "less than everything else" in { - assert(Bottom < fromSeconds(0)) - assert(Bottom < fromNanoseconds(Long.MaxValue)) - assert(Bottom < fromNanoseconds(Long.MinValue)) - } + test("Bottom should become undefined when added with Top or subtracted by bottom") { + assert(Bottom + Duration.Top == Undefined) + assert(Bottom - Duration.Bottom == Undefined) + } - "less than Top" in { - assert(Bottom < Top) - } + test("Bottom should always be min") { + assert((Bottom min Top) == Bottom) + assert((Bottom min fromNanoseconds(0)) == Bottom) + } - "equal to itself" in { - assert(Bottom == Bottom) - } + test("Bottom should less than everything else") { + assert(Bottom < fromSeconds(0)) + assert(Bottom < fromNanoseconds(Long.MaxValue)) + assert(Bottom < fromNanoseconds(Long.MinValue)) + } - "more or less equals only to itself" in { - assert(Bottom.moreOrLessEquals(Bottom, Duration.Top) == true) - assert(Bottom.moreOrLessEquals(Bottom, Duration.Zero) == true) - assert(Bottom.moreOrLessEquals(Top, Duration.Bottom) == false) - assert(Bottom.moreOrLessEquals(Top, Duration.Zero) == false) - assert(Bottom.moreOrLessEquals(fromSeconds(0), Duration.Top) == true) - assert(Bottom.moreOrLessEquals(fromSeconds(0), Duration.Bottom) == false) - } + test("Bottom should less than Top") { + assert(Bottom < Top) + } - "Undefined diff to Bottom" in { - assert((Bottom diff Bottom) == Duration.Undefined) - } + test("Bottom should equal to itself") { + assert(Bottom == Bottom) } - "Undefined" should { - "be impermeable to any arithmetic" in { - assert(Undefined + 0.seconds == Undefined) - assert(Undefined + 100.seconds == Undefined) - assert(Undefined + Duration.fromNanoseconds(Long.MaxValue) == Undefined) - } + test("Bottom should more or less equals only to itself") { + assert(Bottom.moreOrLessEquals(Bottom, Duration.Top) == true) + assert(Bottom.moreOrLessEquals(Bottom, Duration.Zero) == true) + assert(Bottom.moreOrLessEquals(Top, Duration.Bottom) == false) + assert(Bottom.moreOrLessEquals(Top, Duration.Zero) == false) + assert(Bottom.moreOrLessEquals(fromSeconds(0), Duration.Top) == true) + assert(Bottom.moreOrLessEquals(fromSeconds(0), Duration.Bottom) == false) + } - "become undefined when added with Top or subtracted by bottom" in { - assert(Undefined + Duration.Top == Undefined) - assert(Undefined - Duration.Undefined == Undefined) - } + test("Bottom should Undefined diff to Bottom") { + assert((Bottom diff Bottom) == Duration.Undefined) + } - "always be max" in { - assert((Undefined max Top) == Undefined) - assert((Undefined max fromNanoseconds(0)) == Undefined) - } + test("Undefined should be impermeable to any arithmetic") { + assert(Undefined + 0.seconds == Undefined) + assert(Undefined + 100.seconds == Undefined) + assert(Undefined + Duration.fromNanoseconds(Long.MaxValue) == Undefined) + } - "greater than everything else" in { - assert(fromSeconds(0) < Undefined) - assert(Top < Undefined) - assert(fromNanoseconds(Long.MaxValue) < Undefined) - } + test("Undefined should become undefined when added with Top or subtracted by bottom") { + assert(Undefined + Duration.Top == Undefined) + assert(Undefined - Duration.Undefined == Undefined) + } - "equal to itself" in { - assert(Undefined == Undefined) - } + test("Undefined should always be max") { + assert((Undefined max Top) == Undefined) + assert((Undefined max fromNanoseconds(0)) == Undefined) + } - "not more or less equal to anything" in { - assert(Undefined.moreOrLessEquals(Undefined, Duration.Top) == false) - assert(Undefined.moreOrLessEquals(Undefined, Duration.Zero) == false) - assert(Undefined.moreOrLessEquals(Top, Duration.Undefined) == true) - assert(Undefined.moreOrLessEquals(Top, Duration.Zero) == false) - assert(Undefined.moreOrLessEquals(fromSeconds(0), Duration.Top) == false) - assert(Undefined.moreOrLessEquals(fromSeconds(0), Duration.Undefined) == true) - } + test("Undefined should greater than everything else") { + assert(fromSeconds(0) < Undefined) + assert(Top < Undefined) + assert(fromNanoseconds(Long.MaxValue) < Undefined) + } - "Undefined on diff" in { - assert((Undefined diff Top) == Duration.Undefined) - assert((Undefined diff Bottom) == Duration.Undefined) - assert((Undefined diff fromNanoseconds(123)) == Duration.Undefined) - } + test("Undefined should equal to itself") { + assert(Undefined == Undefined) } - "values" should { - "reflect their underlying value" in { - val nss = Seq( - 2592000000000000000L, // 30000.days - 1040403005001003L, // 12.days+1.hour+3.seconds+5.milliseconds+1.microsecond+3.nanoseconds - 123000000000L, // 123.seconds - 1L - ) + test("Undefined should not more or less equal to anything") { + assert(Undefined.moreOrLessEquals(Undefined, Duration.Top) == false) + assert(Undefined.moreOrLessEquals(Undefined, Duration.Zero) == false) + assert(Undefined.moreOrLessEquals(Top, Duration.Undefined) == true) + assert(Undefined.moreOrLessEquals(Top, Duration.Zero) == false) + assert(Undefined.moreOrLessEquals(fromSeconds(0), Duration.Top) == false) + assert(Undefined.moreOrLessEquals(fromSeconds(0), Duration.Undefined) == true) + } - for (ns <- nss) { - val t = fromNanoseconds(ns) - assert(t.inNanoseconds == ns) - assert(t.inMicroseconds == ns / 1000L) - assert(t.inMilliseconds == ns / 1000000L) - assert(t.inLongSeconds == ns / 1000000000L) - assert(t.inMinutes == ns / 60000000000L) - assert(t.inHours == ns / 3600000000000L) - assert(t.inDays == ns / 86400000000000L) - } - } + test("Undefined should Undefined on diff") { + assert((Undefined diff Top) == Duration.Undefined) + assert((Undefined diff Bottom) == Duration.Undefined) + assert((Undefined diff fromNanoseconds(123)) == Duration.Undefined) } - "inSeconds" should { - "equal inLongSeconds when in 32-bit range" in { - val nss = Seq( - 315370851000000000L, // 3650.days+3.hours+51.seconds - 1040403005001003L, // 12.days+1.hour+3.seconds+5.milliseconds+1.microsecond+3.nanoseconds - 1L - ) - for (ns <- nss) { - val t = fromNanoseconds(ns) - assert(t.inLongSeconds == t.inSeconds) - } - } - "clamp value to Int.MinValue or MaxValue when out of range" in { - val longNs = 2160000000000000000L // 25000.days - assert(fromNanoseconds(longNs).inSeconds == Int.MaxValue) - assert(fromNanoseconds(-longNs).inSeconds == Int.MinValue) + test("values should reflect their underlying value") { + val nss = Seq( + 2592000000000000000L, // 30000.days + 1040403005001003L, // 12.days+1.hour+3.seconds+5.milliseconds+1.microsecond+3.nanoseconds + 123000000000L, // 123.seconds + 1L + ) + + for (ns <- nss) { + val t = fromNanoseconds(ns) + assert(t.inNanoseconds == ns) + assert(t.inMicroseconds == ns / 1000L) + assert(t.inMilliseconds == ns / 1000000L) + assert(t.inLongSeconds == ns / 1000000000L) + assert(t.inMinutes == ns / 60000000000L) + assert(t.inHours == ns / 3600000000000L) + assert(t.inDays == ns / 86400000000000L) } } - "rounding" should { - - "maintain top and bottom" in { - assert(Top.floor(1.hour) == Top) - assert(Bottom.floor(1.hour) == Bottom) + test("inSeconds should equal inLongSeconds when in 32-bit range") { + val nss = Seq( + 315370851000000000L, // 3650.days+3.hours+51.seconds + 1040403005001003L, // 12.days+1.hour+3.seconds+5.milliseconds+1.microsecond+3.nanoseconds + 1L + ) + for (ns <- nss) { + val t = fromNanoseconds(ns) + assert(t.inLongSeconds == t.inSeconds) } + } + test("inSeconds should clamp value to Int.MinValue or MaxValue when out of range") { + val longNs = 2160000000000000000L // 25000.days + assert(fromNanoseconds(longNs).inSeconds == Int.MaxValue) + assert(fromNanoseconds(-longNs).inSeconds == Int.MinValue) + } - "divide by zero" in { - assert(Zero.floor(Duration.Zero) == Undefined) - assert(fromSeconds(1).floor(Duration.Zero) == Top) - assert(fromSeconds(-1).floor(Duration.Zero) == Bottom) - } + test("rounding should maintain top and bottom") { + assert(Top.floor(1.hour) == Top) + assert(Bottom.floor(1.hour) == Bottom) + } - "deal with undefineds" in { - assert(Undefined.floor(0.seconds) == Undefined) - assert(Undefined.floor(Duration.Top) == Undefined) - assert(Undefined.floor(Duration.Bottom) == Undefined) - assert(Undefined.floor(Duration.Undefined) == Undefined) - } + test("rounding should divide by zero") { + assert(Zero.floor(Duration.Zero) == Undefined) + assert(fromSeconds(1).floor(Duration.Zero) == Top) + assert(fromSeconds(-1).floor(Duration.Zero) == Bottom) + } - "round to itself" in { - for (s <- Seq(Long.MinValue, -1, 1, Long.MaxValue); t = s.nanoseconds) - assert(t.floor(t.inNanoseconds.nanoseconds) == t) - } + test("rounding should deal with undefineds") { + assert(Undefined.floor(0.seconds) == Undefined) + assert(Undefined.floor(Duration.Top) == Undefined) + assert(Undefined.floor(Duration.Bottom) == Undefined) + assert(Undefined.floor(Duration.Undefined) == Undefined) } - "floor" should { - "round down" in { - assert(60.seconds.floor(1.minute) == 60.seconds) - assert(100.seconds.floor(1.minute) == 60.seconds) - assert(119.seconds.floor(1.minute) == 60.seconds) - assert(120.seconds.floor(1.minute) == 120.seconds) - } + test("rounding should round to itself") { + for (s <- Seq(Long.MinValue, -1, 1, Long.MaxValue); t = s.nanoseconds) + assert(t.floor(t.inNanoseconds.nanoseconds) == t) } - "ceiling" should { - "round up" in { - assert(60.seconds.ceil(1.minute) == 60.seconds) - assert(100.seconds.ceil(1.minute) == 120.seconds) - assert(119.seconds.ceil(1.minute) == 120.seconds) - assert(120.seconds.ceil(1.minute) == 120.seconds) - } + test("floor should round down") { + assert(60.seconds.floor(1.minute) == 60.seconds) + assert(100.seconds.floor(1.minute) == 60.seconds) + assert(119.seconds.floor(1.minute) == 60.seconds) + assert(120.seconds.floor(1.minute) == 120.seconds) } - "from*" should { - "never over/under flow nanos" in { - for (v <- Seq(Long.MinValue, Long.MaxValue)) { - fromNanoseconds(v) match { - case Nanoseconds(ns) => assert(ns == v) - } + test("ceiling should round up") { + assert(60.seconds.ceil(1.minute) == 60.seconds) + assert(100.seconds.ceil(1.minute) == 120.seconds) + assert(119.seconds.ceil(1.minute) == 120.seconds) + assert(120.seconds.ceil(1.minute) == 120.seconds) + } + + test("from* should never over/under flow nanos") { + for (v <- Seq(Long.MinValue, Long.MaxValue)) { + fromNanoseconds(v) match { + case Nanoseconds(ns) => assert(ns == v) } } + } - "overflow millis" in { - val millis = TimeUnit.NANOSECONDS.toMillis(Long.MaxValue) - fromMilliseconds(millis) match { - case Nanoseconds(ns) => assert(ns == millis * 1e6) - } - assert(fromMilliseconds(millis + 1) == Top) + test("from* should overflow millis") { + val millis = TimeUnit.NANOSECONDS.toMillis(Long.MaxValue) + fromMilliseconds(millis) match { + case Nanoseconds(ns) => assert(ns == millis * 1e6) } + assert(fromMilliseconds(millis + 1) == Top) + } - "underflow millis" in { - val millis = TimeUnit.NANOSECONDS.toMillis(Long.MinValue) - fromMilliseconds(millis) match { - case Nanoseconds(ns) => assert(ns == millis * 1e6) - } - assert(fromMilliseconds(millis - 1) == Bottom) + test("from* should underflow millis") { + val millis = TimeUnit.NANOSECONDS.toMillis(Long.MinValue) + fromMilliseconds(millis) match { + case Nanoseconds(ns) => assert(ns == millis * 1e6) } + assert(fromMilliseconds(millis - 1) == Bottom) } } -@RunWith(classOf[JUnitRunner]) -class TimeFormatTest extends WordSpec { - "TimeFormat" should { - "format correctly with non US locale" in { - val locale = Locale.GERMAN - val format = "EEEE" - val timeFormat = new TimeFormat(format, Some(locale)) - val day = "Donnerstag" - assert(timeFormat.parse(day).format(format, locale) == day) - } +class TimeFormatTest extends AnyFunSuite { + test("TimeFormat should format correctly with non US locale") { + val locale = Locale.GERMAN + val format = "EEEE" + val timeFormat = new TimeFormat(format, Some(locale)) + val day = "Donnerstag" + assert(timeFormat.parse(day).format(format, locale) == day) + } - "set UTC timezone as default" in { - val format = "HH:mm" - val timeFormat = new TimeFormat(format) - assert(timeFormat.parse("16:04").format(format) == "16:04") - } + test("TimeFormat should set UTC timezone as default") { + val format = "HH:mm" + val timeFormat = new TimeFormat(format) + assert(timeFormat.parse("16:04").format(format) == "16:04") + } - "set non-UTC timezone correctly" in { - val format = "HH:mm" - val timeFormat = new TimeFormat(format, TimeZone.getTimeZone("EST")) - assert(timeFormat.parse("16:04").format(format) == "21:04") - } + test("TimeFormat should set non-UTC timezone correctly") { + val format = "HH:mm" + val timeFormat = new TimeFormat(format, TimeZone.getTimeZone("EST")) + assert(timeFormat.parse("16:04").format(format) == "21:04") } } -@RunWith(classOf[JUnitRunner]) -class TimeTest extends { val ops = Time } with TimeLikeSpec[Time] with Eventually -with IntegrationPatience { - - "Time" should { - "work in collections" in { - val t0 = Time.fromSeconds(100) - val t1 = Time.fromSeconds(100) - assert(t0 == t1) - assert(t0.hashCode == t1.hashCode) - val pairs = List((t0, "foo"), (t1, "bar")) - assert(pairs.groupBy { case (time: Time, value: String) => time } == Map(t0 -> pairs)) - } +class TimeFormatterTest extends AnyFunSuite { + def assertParsesTo( + fmt: String, + str: String, + year: Int, + month: Int, + day: Int, + hour: Int, + minute: Int, + second: Int + ) = { + assert( + TimeFormatter(fmt).parse(str) == Time.fromInstant( + Instant.from(OffsetDateTime.of(year, month, day, hour, minute, second, 0, ZoneOffset.UTC)))) + + } + + test("format correctly with non US locale") { + val locale = Locale.GERMAN + val format = "EEEE" + val timeFormat = TimeFormatter(format, locale) + val day = "Donnerstag" + assert(timeFormat.format(timeFormat.parse(day)) == day) + } - "now should be now" in { - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + test("set UTC timezone as default") { + val format = "HH:mm" + val timeFormat = TimeFormatter(format) + val utcFormat = TimeFormatter(format, zoneId = ZoneOffset.UTC) + assert(timeFormat.format(utcFormat.parse("16:04")) == "16:04") + } + + test("set non-UTC timezone correctly") { + val format = "HH:mm" + val estFormat = TimeFormatter(format, zoneId = ZoneId.of("America/New_York")) + val utcFormat = TimeFormatter(format, zoneId = ZoneOffset.UTC) + assert(utcFormat.format(estFormat.parse("16:04")) == "21:04") + } + + test("parse dates with explicit offset timezone") { + // Known bug in Java 1.8.0 until at least _345, see https://bugs.openjdk.org/browse/JDK-8294853 + // We're conditionally executing this test because pants still runs tests with Java 8. + // Alas, java.lang.Runtime.Version also doesn't exist until Java 9, so we're stuck with this + // somewhat pedestrian way of checking the runtime version... + if (!System.getProperty("java.version").startsWith("1.8.0")) { + assertParsesTo("yyyy-MM-dd HH:mm:ss Z", "2022-09-29 13:33:37 +0100", 2022, 9, 29, 12, 33, 37) } + } - "withTimeAt" in { - val t0 = new Time(123456789L) - Time.withTimeAt(t0) { _ => - assert(Time.now == t0) - Thread.sleep(50) - assert(Time.now == t0) - } - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + test("parse dates with explicit named timezone") { + assertParsesTo("yyyy-MM-dd HH:mm:ss z", "2022-09-29 13:33:37 EST", 2022, 9, 29, 17, 33, 37) + } + + test("parse dates with no timezone") { + assertParsesTo("yyyy-MM-dd HH:mm:ss", "2022-09-29 13:33:37", 2022, 9, 29, 13, 33, 37) + } + + test("parse dates only") { + assertParsesTo("yyyy-MM-dd", "2022-09-29", 2022, 9, 29, 0, 0, 0) + } + + test("parse year-month only") { + assertParsesTo("yyyy-MM", "2022-09", 2022, 9, 1, 0, 0, 0) + } + + test("parse month-day only") { + assertParsesTo("MM-dd", "09-29", 1970, 9, 29, 0, 0, 0) + } + + test("parse hour-minute only") { + assertParsesTo("HH-mm", "13-37", 1970, 1, 1, 13, 37, 0) + } + + test("parse hour only") { + assertParsesTo("HH", "13", 1970, 1, 1, 13, 0, 0) + } +} + +trait TimeOps { val ops: Time.type = Time } +class TimeTest + extends AnyFunSuite + with TimeOps + with TimeLikeSpec[Time] + with Eventually + with IntegrationPatience { + + test("Time should work in collections") { + val t0 = Time.fromSeconds(100) + val t1 = Time.fromSeconds(100) + assert(t0 == t1) + assert(t0.hashCode == t1.hashCode) + val pairs = List((t0, "foo"), (t1, "bar")) + assert(pairs.groupBy { case (time: Time, value: String) => time } == Map(t0 -> pairs)) + } + + test("Time now should be now") { + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } + + test("Time withTimeAt") { + val t0 = new Time(123456789L) + Time.withTimeAt(t0) { _ => + assert(Time.now == t0) + Thread.sleep(50) + assert(Time.now == t0) } + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } - "withTimeAt nested" in { - val t0 = new Time(123456789L) - val t1 = t0 + 10.minutes - Time.withTimeAt(t0) { _ => - assert(Time.now == t0) - Time.withTimeAt(t1) { _ => - assert(Time.now == t1) - } - assert(Time.now == t0) - } - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + test("Time withTimeAt nested") { + val t0 = new Time(123456789L) + val t1 = t0 + 10.minutes + Time.withTimeAt(t0) { _ => + assert(Time.now == t0) + Time.withTimeAt(t1) { _ => assert(Time.now == t1) } + assert(Time.now == t0) } + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } - "withTimeAt threaded" in { - val t0 = new Time(314159L) - val t1 = new Time(314160L) - Time.withTimeAt(t0) { tc => - assert(Time.now == t0) - Thread.sleep(50) - assert(Time.now == t0) - tc.advance(Duration.fromNanoseconds(1)) - assert(Time.now == t1) - tc.set(t0) - assert(Time.now == t0) - @volatile var threadTime: Option[Time] = None - val thread = new Thread { - override def run(): Unit = { - threadTime = Some(Time.now) - } + test("Time withTimeAt threaded") { + val t0 = new Time(314159L) + val t1 = new Time(314160L) + Time.withTimeAt(t0) { tc => + assert(Time.now == t0) + Thread.sleep(50) + assert(Time.now == t0) + tc.advance(Duration.fromNanoseconds(1)) + assert(Time.now == t1) + tc.set(t0) + assert(Time.now == t0) + @volatile var threadTime: Option[Time] = None + val thread = new Thread { + override def run(): Unit = { + threadTime = Some(Time.now) } - thread.start() - thread.join() - assert(threadTime.get != t0) } - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + thread.start() + thread.join() + assert(threadTime.get != t0) } + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } - "withTimeFunction" in { - val t0 = Time.now - var t = t0 - Time.withTimeFunction(t) { _ => - assert(Time.now == t0) - Thread.sleep(50) - assert(Time.now == t0) - val delta = 100.milliseconds - t += delta - assert(Time.now == t0 + delta) - } + test("Time withTimeFunction") { + val t0 = Time.now + var t = t0 + Time.withTimeFunction(t) { _ => + assert(Time.now == t0) + Thread.sleep(50) + assert(Time.now == t0) + val delta = 100.milliseconds + t += delta + assert(Time.now == t0 + delta) } + } - "withCurrentTimeFrozen" in { - val t0 = new Time(123456789L) - Time.withCurrentTimeFrozen { _ => - val t0 = Time.now - Thread.sleep(50) - assert(Time.now == t0) - } - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + test("Time withCurrentTimeFrozen") { + val t0 = new Time(123456789L) + Time.withCurrentTimeFrozen { _ => + val t0 = Time.now + Thread.sleep(50) + assert(Time.now == t0) } + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } - "advance" in { - val t0 = new Time(123456789L) - val delta = 5.seconds - Time.withTimeAt(t0) { tc => - assert(Time.now == t0) - tc.advance(delta) - assert(Time.now == (t0 + delta)) - } - assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + test("Time advance") { + val t0 = new Time(123456789L) + val delta = 5.seconds + Time.withTimeAt(t0) { tc => + assert(Time.now == t0) + tc.advance(delta) + assert(Time.now == (t0 + delta)) } + assert((Time.now.inMillis - System.currentTimeMillis).abs < 20L) + } - "sleep" in { - Time.withCurrentTimeFrozen { ctl => - val ctx = Local.save() - val r = new Runnable { - def run(): Unit = { - Local.restore(ctx) - Time.sleep(5.seconds) - } + test("Time sleep") { + Time.withCurrentTimeFrozen { ctl => + val ctx = Local.save() + val r = new Runnable { + def run(): Unit = { + Local.restore(ctx) + Time.sleep(5.seconds) } - @volatile var x = 0 - val t = new Thread(r) - t.start() - assert(t.isAlive == true) - eventually { assert(t.getState == Thread.State.TIMED_WAITING) } - ctl.advance(5.seconds) - t.join() - assert(t.isAlive == false) } + @volatile var x = 0 + val t = new Thread(r) + t.start() + assert(t.isAlive == true) + eventually { + assert(t.getState == Thread.State.TIMED_WAITING) + } + ctl.advance(5.seconds) + t.join() + assert(t.isAlive == false) } + } - "compare" in { - assert(10.seconds.afterEpoch < 11.seconds.afterEpoch) - assert(10.seconds.afterEpoch == 10.seconds.afterEpoch) - assert(11.seconds.afterEpoch > 10.seconds.afterEpoch) - assert(Time.fromMilliseconds(Long.MaxValue) > Time.now) - } + test("Time compare") { + assert(10.seconds.afterEpoch < 11.seconds.afterEpoch) + assert(10.seconds.afterEpoch == 10.seconds.afterEpoch) + assert(11.seconds.afterEpoch > 10.seconds.afterEpoch) + assert(Time.fromMilliseconds(Long.MaxValue) > Time.now) + } - "equals" in { - assert(Time.Top == Time.Top) - assert(Time.Top != Time.Bottom) - assert(Time.Top != Time.Undefined) + test("Time equals") { + assert(Time.Top == Time.Top) + assert(Time.Top != Time.Bottom) + assert(Time.Top != Time.Undefined) - assert(Time.Bottom != Time.Top) - assert(Time.Bottom == Time.Bottom) - assert(Time.Bottom != Time.Undefined) + assert(Time.Bottom != Time.Top) + assert(Time.Bottom == Time.Bottom) + assert(Time.Bottom != Time.Undefined) - assert(Time.Undefined != Time.Top) - assert(Time.Undefined != Time.Bottom) - assert(Time.Undefined == Time.Undefined) + assert(Time.Undefined != Time.Top) + assert(Time.Undefined != Time.Bottom) + assert(Time.Undefined == Time.Undefined) - val now = Time.now - assert(now == now) - assert(now == Time.fromNanoseconds(now.inNanoseconds)) - assert(now != now + 1.nanosecond) - } + val now = Time.now + assert(now == now) + assert(now == Time.fromNanoseconds(now.inNanoseconds)) + assert(now != now + 1.nanosecond) + } - "+ delta" in { - assert(10.seconds.afterEpoch + 5.seconds == 15.seconds.afterEpoch) - } + test("Time + delta") { + assert(10.seconds.afterEpoch + 5.seconds == 15.seconds.afterEpoch) + } - "- delta" in { - assert(10.seconds.afterEpoch - 5.seconds == 5.seconds.afterEpoch) - } + test("Time - delta") { + assert(10.seconds.afterEpoch - 5.seconds == 5.seconds.afterEpoch) + } - "- time" in { - assert(10.seconds.afterEpoch - 5.seconds.afterEpoch == 5.seconds) - } + test("Time - time") { + assert(10.seconds.afterEpoch - 5.seconds.afterEpoch == 5.seconds) + } - "max" in { - assert((10.seconds.afterEpoch max 5.seconds.afterEpoch) == 10.seconds.afterEpoch) - assert((5.seconds.afterEpoch max 10.seconds.afterEpoch) == 10.seconds.afterEpoch) - } + test("Time max") { + assert((10.seconds.afterEpoch max 5.seconds.afterEpoch) == 10.seconds.afterEpoch) + assert((5.seconds.afterEpoch max 10.seconds.afterEpoch) == 10.seconds.afterEpoch) + } - "min" in { - assert((10.seconds.afterEpoch min 5.seconds.afterEpoch) == 5.seconds.afterEpoch) - assert((5.seconds.afterEpoch min 10.seconds.afterEpoch) == 5.seconds.afterEpoch) - } + test("Time min") { + assert((10.seconds.afterEpoch min 5.seconds.afterEpoch) == 5.seconds.afterEpoch) + assert((5.seconds.afterEpoch min 10.seconds.afterEpoch) == 5.seconds.afterEpoch) + } - "moreOrLessEquals" in { - val now = Time.now - assert(now.moreOrLessEquals(now + 1.second, 1.second) == true) - assert(now.moreOrLessEquals(now - 1.seconds, 1.second) == true) - assert(now.moreOrLessEquals(now + 2.seconds, 1.second) == false) - assert(now.moreOrLessEquals(now - 2.seconds, 1.second) == false) - } + test("Time moreOrLessEquals") { + val now = Time.now + assert(now.moreOrLessEquals(now + 1.second, 1.second) == true) + assert(now.moreOrLessEquals(now - 1.seconds, 1.second) == true) + assert(now.moreOrLessEquals(now + 2.seconds, 1.second) == false) + assert(now.moreOrLessEquals(now - 2.seconds, 1.second) == false) + } - "floor" in { - val format = new TimeFormat("yyyy-MM-dd HH:mm:ss.SSS") - val t0 = format.parse("2010-12-24 11:04:07.567") - assert(t0.floor(1.millisecond) == t0) - assert(t0.floor(10.milliseconds) == format.parse("2010-12-24 11:04:07.560")) - assert(t0.floor(1.second) == format.parse("2010-12-24 11:04:07.000")) - assert(t0.floor(5.second) == format.parse("2010-12-24 11:04:05.000")) - assert(t0.floor(1.minute) == format.parse("2010-12-24 11:04:00.000")) - assert(t0.floor(1.hour) == format.parse("2010-12-24 11:00:00.000")) - } + test("Time floor") { + val format = new TimeFormat("yyyy-MM-dd HH:mm:ss.SSS") + val t0 = format.parse("2010-12-24 11:04:07.567") + assert(t0.floor(1.millisecond) == t0) + assert(t0.floor(10.milliseconds) == format.parse("2010-12-24 11:04:07.560")) + assert(t0.floor(1.second) == format.parse("2010-12-24 11:04:07.000")) + assert(t0.floor(5.second) == format.parse("2010-12-24 11:04:05.000")) + assert(t0.floor(1.minute) == format.parse("2010-12-24 11:04:00.000")) + assert(t0.floor(1.hour) == format.parse("2010-12-24 11:00:00.000")) + } - "since" in { - val t0 = Time.now - val t1 = t0 + 10.seconds - assert(t1.since(t0) == 10.seconds) - assert(t0.since(t1) == (-10).seconds) - } + test("Time since") { + val t0 = Time.now + val t1 = t0 + 10.seconds + assert(t1.since(t0) == 10.seconds) + assert(t0.since(t1) == (-10).seconds) + } - "sinceEpoch" in { - val t0 = Time.epoch + 100.hours - assert(t0.sinceEpoch == 100.hours) - } + test("Time sinceEpoch") { + val t0 = Time.epoch + 100.hours + assert(t0.sinceEpoch == 100.hours) + } - "sinceNow" in { - Time.withCurrentTimeFrozen { _ => - val t0 = Time.now + 100.hours - assert(t0.sinceNow == 100.hours) - } + test("Time sinceNow") { + Time.withCurrentTimeFrozen { _ => + val t0 = Time.now + 100.hours + assert(t0.sinceNow == 100.hours) } + } - "fromSeconds(Double)" in { - val tolerance = 2.microseconds // we permit 1us slop - - forAll { i: Int => - assert( - Time.fromSeconds(i).moreOrLessEquals(Time.fromFractionalSeconds(i.toDouble), tolerance) - ) - } + test("Time fromFractionalSeconds") { + val tolerance = 2.microseconds // we permit 1us slop - forAll { d: Double => - val magic = 9223372036854775L // cribbed from Time.fromMicroseconds - val microseconds = d * 1.second.inMicroseconds - whenever(microseconds > -magic && microseconds < magic) { - assert( - Time - .fromMicroseconds(microseconds.toLong) - .moreOrLessEquals(Time.fromFractionalSeconds(d), tolerance) - ) - } - } + forAll { (i: Int) => + assert( + Time.fromSeconds(i).moreOrLessEquals(Time.fromFractionalSeconds(i.toDouble), tolerance) + ) + } - forAll { l: Long => - val seconds: Double = l.toDouble / 1.second.inNanoseconds + forAll { (d: Double) => + val magic = 9223372036854775L // cribbed from Time.fromMicroseconds + val microseconds = d * 1.second.inMicroseconds + whenever(microseconds > -magic && microseconds < magic) { assert( - Time.fromFractionalSeconds(seconds).moreOrLessEquals(Time.fromNanoseconds(l), tolerance) + Time + .fromMicroseconds(microseconds.toLong) + .moreOrLessEquals(Time.fromFractionalSeconds(d), tolerance) ) } } - "fromMicroseconds" in { - assert(Time.fromMicroseconds(0).inNanoseconds == 0L) - assert(Time.fromMicroseconds(-1).inNanoseconds == -1L * 1000L) - - assert(Time.fromMicroseconds(Long.MaxValue).inNanoseconds == Long.MaxValue) - assert(Time.fromMicroseconds(Long.MaxValue - 1) == Time.Top) - - assert(Time.fromMicroseconds(Long.MinValue) == Time.Bottom) - assert(Time.fromMicroseconds(Long.MinValue + 1) == Time.Bottom) - - val currentTimeMicros = System.currentTimeMillis() * 1000 + forAll { (l: Long) => + val seconds: Double = l.toDouble / 1.second.inNanoseconds assert( - Time - .fromMicroseconds(currentTimeMicros) - .inNanoseconds == currentTimeMicros.microseconds.inNanoseconds + Time.fromFractionalSeconds(seconds).moreOrLessEquals(Time.fromNanoseconds(l), tolerance) ) } + } - "fromMillis" in { - assert(Time.fromMilliseconds(0).inNanoseconds == 0L) - assert(Time.fromMilliseconds(-1).inNanoseconds == -1L * 1000000L) + test("Time fromMicroseconds") { + assert(Time.fromMicroseconds(0).inNanoseconds == 0L) + assert(Time.fromMicroseconds(-1).inNanoseconds == -1L * 1000L) - assert(Time.fromMilliseconds(Long.MaxValue).inNanoseconds == Long.MaxValue) - assert(Time.fromMilliseconds(Long.MaxValue - 1) == Time.Top) + assert(Time.fromMicroseconds(Long.MaxValue).inNanoseconds == Long.MaxValue) + assert(Time.fromMicroseconds(Long.MaxValue - 1) == Time.Top) - assert(Time.fromMilliseconds(Long.MinValue) == Time.Bottom) - assert(Time.fromMilliseconds(Long.MinValue + 1) == Time.Bottom) + assert(Time.fromMicroseconds(Long.MinValue) == Time.Bottom) + assert(Time.fromMicroseconds(Long.MinValue + 1) == Time.Bottom) - val currentTimeMs = System.currentTimeMillis - assert(Time.fromMilliseconds(currentTimeMs).inNanoseconds == currentTimeMs * 1000000L) + val currentTimeMicros = System.currentTimeMillis() * 1000 + assert( + Time + .fromMicroseconds(currentTimeMicros) + .inNanoseconds == currentTimeMicros.microseconds.inNanoseconds + ) + } + + test("Time fromMillis") { + assert(Time.fromMilliseconds(0).inNanoseconds == 0L) + assert(Time.fromMilliseconds(-1).inNanoseconds == -1L * 1000000L) + + assert(Time.fromMilliseconds(Long.MaxValue).inNanoseconds == Long.MaxValue) + assert(Time.fromMilliseconds(Long.MaxValue - 1) == Time.Top) + + assert(Time.fromMilliseconds(Long.MinValue) == Time.Bottom) + assert(Time.fromMilliseconds(Long.MinValue + 1) == Time.Bottom) + + val currentTimeMs = System.currentTimeMillis + assert(Time.fromMilliseconds(currentTimeMs).inNanoseconds == currentTimeMs * 1000000L) + } + + test("Time fromMinutes") { + assert(Time.fromMinutes(0).inNanoseconds == 0L) + assert(Time.fromMinutes(-1).inNanoseconds == -60L * 1000000000L) + + assert(Time.fromMinutes(Int.MaxValue).inNanoseconds == Long.MaxValue) + assert(Time.fromMinutes(Int.MaxValue) == Time.Top) + assert(Time.fromMinutes(Int.MinValue) == Time.Bottom) + } + + test("Time fromHours") { + assert(Time.fromHours(1).inNanoseconds == Time.fromMinutes(60).inNanoseconds) + + assert(Time.fromHours(0).inNanoseconds == 0L) + assert(Time.fromHours(-1).inNanoseconds == -3600L * 1000000000L) + + assert(Time.fromHours(Int.MaxValue).inNanoseconds == Long.MaxValue) + assert(Time.fromHours(Int.MaxValue) == Time.Top) + assert(Time.fromHours(Int.MinValue) == Time.Bottom) + } + + test("Time fromDays") { + assert(Time.fromDays(1).inNanoseconds == Time.fromHours(24).inNanoseconds) + + assert(Time.fromDays(0).inNanoseconds == 0L) + assert(Time.fromDays(-1).inNanoseconds == -3600L * 24L * 1000000000L) + + assert(Time.fromDays(Int.MaxValue).inNanoseconds == Long.MaxValue) + assert(Time.fromDays(Int.MaxValue) == Time.Top) + assert(Time.fromDays(Int.MinValue) == Time.Bottom) + } + + test("Time until") { + val t0 = Time.now + val t1 = t0 + 10.seconds + assert(t0.until(t1) == 10.seconds) + assert(t1.until(t0) == (-10).seconds) + } + + test("Time untilEpoch") { + val t0 = Time.epoch - 100.hours + assert(t0.untilEpoch == 100.hours) + } + + test("Time untilNow") { + Time.withCurrentTimeFrozen { _ => + val t0 = Time.now - 100.hours + assert(t0.untilNow == 100.hours) } + } - "until" in { - val t0 = Time.now - val t1 = t0 + 10.seconds - assert(t0.until(t1) == 10.seconds) - assert(t1.until(t0) == (-10).seconds) + test("Time toInstant") { + Time.withCurrentTimeFrozen { _ => + val instant = Instant.ofEpochMilli(Time.now.inMilliseconds) + assert(instant.toEpochMilli == Time.now.inMillis) + // java.time.Instant:getNano returns the nanoseconds of the second + assert(instant.getNano == Time.now.inNanoseconds % Duration.NanosPerSecond) } + } - "untilEpoch" in { - val t0 = Time.epoch - 100.hours - assert(t0.untilEpoch == 100.hours) + test("Time toZonedDateTime") { + Time.withCurrentTimeFrozen { _ => + val zonedDateTime = + ZonedDateTime.ofInstant(Instant.ofEpochMilli(Time.now.inMilliseconds), ZoneOffset.UTC) + assert(Time.now.toZonedDateTime == zonedDateTime) + // java.time.Instant:getNano returns the nanoseconds of the second + assert(zonedDateTime.getNano == Time.now.inNanoseconds % Duration.NanosPerSecond) } + } - "untilNow" in { - Time.withCurrentTimeFrozen { _ => - val t0 = Time.now - 100.hours - assert(t0.untilNow == 100.hours) - } + test("Time toOffsetDateTime") { + Time.withCurrentTimeFrozen { _ => + val offsetDateTime = + OffsetDateTime.ofInstant(Instant.ofEpochMilli(Time.now.inMilliseconds), ZoneOffset.UTC) + assert(Time.now.toOffsetDateTime == offsetDateTime) + // java.time.Instant:getNano returns the nanoseconds of the second + assert(offsetDateTime.getNano == Time.now.inNanoseconds % Duration.NanosPerSecond) } } + + test("at") { + assert(Time.at("1970-01-01 00:00:00 +0100").inLongSeconds == -3600) + assert(Time.at("1970-01-01 00:00:00 +0000").inLongSeconds == 0) + assert(Time.at("1970-01-01 00:00:00 UTC").inLongSeconds == 0) + assert(Time.at("1970-01-01 00:00:00 GMT").inLongSeconds == 0) + assert(Time.at("1970-01-01 00:00:00 EST").inLongSeconds == 18000) + } + + test("fromRss") { + assert(Time.fromRss("Thu, 1 Jan 1970 00:00:00 +0100").inLongSeconds == -3600) + assert(Time.fromRss("Thu, 1 Jan 1970 00:00:00 +0000").inLongSeconds == 0) + assert(Time.fromRss("Thu, 1 Jan 1970 00:00:00 UTC").inLongSeconds == 0) + assert(Time.fromRss("Thu, 1 Jan 1970 00:00:00 GMT").inLongSeconds == 0) + assert(Time.fromRss("Thu, 1 Jan 1970 00:00:00 EST").inLongSeconds == 18000) + } } diff --git a/util-core/src/test/scala/com/twitter/util/TimerTest.scala b/util-core/src/test/scala/com/twitter/util/TimerTest.scala index 18c2f594ac..6954a1190f 100644 --- a/util-core/src/test/scala/com/twitter/util/TimerTest.scala +++ b/util-core/src/test/scala/com/twitter/util/TimerTest.scala @@ -1,19 +1,16 @@ package com.twitter.util -import com.twitter.conversions.time._ -import java.util.concurrent.{CancellationException, ExecutorService} +import com.twitter.conversions.DurationOps._ +import java.util.concurrent.{CancellationException, ExecutorService, TimeUnit, CountDownLatch} import java.util.concurrent.atomic.AtomicInteger -import org.junit.runner.RunWith import org.mockito.ArgumentCaptor -import org.mockito.Matchers.any +import org.mockito.ArgumentMatchers.any import org.mockito.Mockito.{never, verify, when} -import org.scalatest.FunSuite import org.scalatest.concurrent.{IntegrationPatience, Eventually} -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class TimerTest extends FunSuite with MockitoSugar with Eventually with IntegrationPatience { +class TimerTest extends AnyFunSuite with MockitoSugar with Eventually with IntegrationPatience { private def testTimerRunsWithLocals(timer: Timer): Unit = { val timerLocal = new AtomicInteger(0) @@ -140,7 +137,7 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat throw new scala.MatchError("huh") } - latch.await(30.seconds) + latch.await(30, TimeUnit.SECONDS) assert(errors == 1) @@ -151,7 +148,7 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat latch.countDown() } - latch.await(30.seconds) + latch.await(30, TimeUnit.SECONDS) assert(result == 2) assert(errors == 1) @@ -289,7 +286,7 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat Time.withCurrentTimeFrozen { ctl => val timer = new MockTimer assert(timer.tasks.isEmpty) - val task = timer.schedule(Time.now + 3.seconds)() + val task = timer.schedule(Time.now + 3.seconds)(()) assert(timer.tasks.size == 1) task.cancel() assert(timer.tasks.isEmpty) @@ -300,7 +297,7 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat Time.withCurrentTimeFrozen { ctl => val timer = new MockTimer assert(timer.tasks.isEmpty) - val task = timer.schedule(3.seconds)() + val task = timer.schedule(3.seconds)(()) eventually { assert(timer.tasks.size == 1) } ctl.advance(30.seconds) timer.tick() @@ -336,10 +333,7 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat } } - private def mockTimerLocalPropagation( - timer: MockTimer, - localValue: Int - ): Int = { + private def mockTimerLocalPropagation(timer: MockTimer, localValue: Int): Int = { Time.withCurrentTimeFrozen { tc => val timerLocal = new AtomicInteger(0) val local = new Local[Int] @@ -359,9 +353,16 @@ class TimerTest extends FunSuite with MockitoSugar with Eventually with Integrat assert(mockTimerLocalPropagation(timer, 99) == 99) } - private class SomeEx extends Exception - private def testTimerUsesLocalMonitor(timer: Timer): Unit = { + // Let's define this exception within this private method since + // the Scala 3.0.1 compiler removes the accessor method to the + // outer instance, `TimerTest`. This is needed when pattern matching + // on the exception in `Monitor.mk`. + // + // See CSL-11237 or https://github.com/lampepfl/dotty/issues/13096. + // The next Scala 3 release will fix this bug. + class SomeEx extends Exception + val seen = new AtomicInteger(0) val monitor = Monitor.mk { case _: SomeEx => diff --git a/util-core/src/test/scala/com/twitter/util/TokenBucketTest.scala b/util-core/src/test/scala/com/twitter/util/TokenBucketTest.scala index 6b36aa46e2..ecfaf91f54 100644 --- a/util-core/src/test/scala/com/twitter/util/TokenBucketTest.scala +++ b/util-core/src/test/scala/com/twitter/util/TokenBucketTest.scala @@ -1,12 +1,9 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class TokenBucketTest extends FunSuite { +class TokenBucketTest extends AnyFunSuite { test("a leaky bucket is leaky") { Time.withCurrentTimeFrozen { tc => val b = TokenBucket.newLeakyBucket(3.seconds, 0, Stopwatch.timeMillis) diff --git a/util-core/src/test/scala/com/twitter/util/TryTest.scala b/util-core/src/test/scala/com/twitter/util/TryTest.scala index 2317b059d4..8c54d5fa58 100644 --- a/util-core/src/test/scala/com/twitter/util/TryTest.scala +++ b/util-core/src/test/scala/com/twitter/util/TryTest.scala @@ -1,13 +1,10 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class TryTest extends FunSuite { +class TryTest extends AnyFunSuite { class MyException extends Exception - val e = new Exception("this is an exception") + val e: Exception = new Exception("this is an exception") test("Try.apply(): should catch exceptions and lift into the Try type") { assert(Try[Int](1) == Return(1)) @@ -214,7 +211,9 @@ class TryTest extends FunSuite { } test("Try to scala.util.Try works") { - import scala.util.{Try => STry, Failure, Success} + import scala.util.{Try => STry} + import scala.util.Failure + import scala.util.Success assert(STry(1) == Try(1).asScala) assert(Try(sys.error("boom!")).asScala match { diff --git a/util-core/src/test/scala/com/twitter/util/TwitterDateFormatTest.scala b/util-core/src/test/scala/com/twitter/util/TwitterDateFormatTest.scala index 07293c9f81..c3fccc5088 100644 --- a/util-core/src/test/scala/com/twitter/util/TwitterDateFormatTest.scala +++ b/util-core/src/test/scala/com/twitter/util/TwitterDateFormatTest.scala @@ -1,47 +1,23 @@ package com.twitter.util import java.util.Locale +import org.scalatest.funsuite.AnyFunSuite -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class TwitterDateFormatTest extends WordSpec { - "TwitterDateFormat" should { - "disallow Y without w" in { - intercept[IllegalArgumentException] { - TwitterDateFormat("YYYYMMDD") - } - intercept[IllegalArgumentException] { - TwitterDateFormat("YMD", Locale.GERMAN) - } - } - - "allow Y with w" in { - TwitterDateFormat("YYYYww") +class TwitterDateFormatTest extends AnyFunSuite { + test("TwitterDateFormat should disallow Y without w") { + intercept[IllegalArgumentException] { + TwitterDateFormat("YYYYMMDD") } - - "allow Y when quoted" in { - TwitterDateFormat("yyyy 'Year'") + intercept[IllegalArgumentException] { + TwitterDateFormat("YMD", Locale.GERMAN) } + } - "stripSingleQuoted" in { - import TwitterDateFormat._ - assert(stripSingleQuoted("") == "") - assert(stripSingleQuoted("YYYY") == "YYYY") - assert(stripSingleQuoted("''") == "") - assert(stripSingleQuoted("'abc'") == "") - assert(stripSingleQuoted("x'abc'") == "x") - assert(stripSingleQuoted("'abc'x") == "x") - assert(stripSingleQuoted("'abc'def'ghi'") == "def") + test("TwitterDateFormat should allow Y with w") { + TwitterDateFormat("YYYYww") + } - intercept[IllegalArgumentException] { - stripSingleQuoted("'abc") - } - intercept[IllegalArgumentException] { - stripSingleQuoted("'abc'def'ghi") - } - } + test("TwitterDateFormat should allow Y when quoted") { + TwitterDateFormat("yyyy 'Year'") } } diff --git a/util-core/src/test/scala/com/twitter/util/U64Test.scala b/util-core/src/test/scala/com/twitter/util/U64Test.scala deleted file mode 100644 index c3e0961415..0000000000 --- a/util-core/src/test/scala/com/twitter/util/U64Test.scala +++ /dev/null @@ -1,155 +0,0 @@ -package com.twitter.util - -import scala.util.Random - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -@RunWith(classOf[JUnitRunner]) -class U64Test extends WordSpec { - import U64._ - - "comparable" in { - { - val a = 0x0000000000000001L - assert(a == 1) - val b = 0x0000000000000002L - assert(b == 2) - - assert(a.u64_<(b) == true) - assert(b.u64_<(a) == false) - } - - { - val a = 0xFFFFFFFFFFFFFFFFL - assert(a == -1) - val b = 0xFFFFFFFFFFFFFFFEL - assert(b == -2) - - assert(a.u64_<(b) == false) - assert(b.u64_<(a) == true) - } - - { - val a = 0xFFFFFFFFFFFFFFFFL - assert(a == -1) - val b = 0x0000000000000001L - assert(b == 1) - - assert(a.u64_<(b) == false) - assert(b.u64_<(a) == true) - } - } - - "comparable in range" in { - assert(0L.u64_within(0, 1) == false) - assert(0L.u64_contained(0, 1) == true) - - // (inverted range) - assert(0L.u64_within(-1, 1) == false) - assert(1L.u64_within(-1, 1) == false) - assert(2L.u64_within(-1, 1) == false) - - assert(0xFFFFFFFFFFFFFFFEL.u64_within(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == true) - assert(0xFFFFFFFFFFFFFFFDL.u64_within(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == false) - assert(0xFFFFFFFFFFFFFFFFL.u64_within(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == false) - - assert(0xFFFFFFFFFFFFFFFEL.u64_contained(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == true) - assert(0xFFFFFFFFFFFFFFFDL.u64_contained(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == true) - assert(0xFFFFFFFFFFFFFFFFL.u64_contained(0xFFFFFFFFFFFFFFFDL, 0xFFFFFFFFFFFFFFFFL) == true) - - // Bit flip area! - assert(0x7FFFFFFFFFFFFFFFL.u64_within(0x7FFFFFFFFFFFFFFFL, 0x8000000000000000L) == false) - assert(0x8000000000000000L.u64_within(0x7FFFFFFFFFFFFFFFL, 0x8000000000000000L) == false) - - assert(0x7FFFFFFFFFFFFFFFL.u64_contained(0x7FFFFFFFFFFFFFFFL, 0x8000000000000000L) == true) - assert(0x8000000000000000L.u64_contained(0x7FFFFFFFFFFFFFFFL, 0x8000000000000000L) == true) - - assert(0x7FFFFFFFFFFFFFFAL.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == false) - assert(0x7FFFFFFFFFFFFFFBL.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == true) - assert(0x7FFFFFFFFFFFFFFFL.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == true) - assert(0x8000000000000000L.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == true) - assert(0x8000000000000001L.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == true) - assert(0x8000000000000009L.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == true) - assert(0x800000000000000AL.u64_within(0x7FFFFFFFFFFFFFFAL, 0x800000000000000AL) == false) - } - - "divisible" in { - assert(10L.u64_/(5L) == 2L) - assert(0xFFFFFFFFFFFFFFFFL / 0x0FFFFFFFFFFFFFFFL == 0L) - assert(0xFFFFFFFFFFFFFFFFL.u64_/(0x0FFFFFFFFFFFFFFFL) == 16L) - - assert(0x7FFFFFFFFFFFFFFFL.u64_/(2) == 0x3FFFFFFFFFFFFFFFL) - - assert(0x8000000000000000L / 2 == 0xc000000000000000L) - assert(0x8000000000000000L.u64_/(2) == 0x4000000000000000L) - - assert(0x8000000000000000L.u64_/(0x8000000000000000L) == 1) - assert(0x8000000000000000L.u64_/(0x8000000000000001L) == 0) - - assert(0xFF00000000000000L.u64_/(0x0F00000000000000L) == 0x11L) - assert(0x8F00000000000000L.u64_/(0x100) == 0x008F000000000000L) - assert(0x8000000000000000L.u64_/(3) == 0x2AAAAAAAAAAAAAAAL) - } - - "ids" should { - "survive conversions" in { - val rng = new Random - - (0 until 10000).foreach { _ => - val id = rng.nextLong - assert(id == (id.toU64ByteArray.toU64Long)) - assert(id == (id.toU64ByteArray.toU64HexString.toU64ByteArray.toU64Long)) - } - } - - "be serializable" in { - assert(0L.toU64HexString == "0000000000000000") - assert(0x0102030405060700L.toU64HexString == "0102030405060700") - assert(0xFFF1F2F3F4F5F6F7L.toU64HexString == "fff1f2f3f4f5f6f7") - } - - "convert from short hex string" in { - assert(new RichU64String("7b").toU64Long == 123L) - } - - "don't silently truncate" in { - intercept[NumberFormatException] { - new RichU64String("318528893302738945") - } - } - - "not parse with +" in { - intercept[NumberFormatException] { "+0".toU64Long } - intercept[NumberFormatException] { "0+".toU64Long } - intercept[NumberFormatException] { "00+0".toU64Long } - intercept[NumberFormatException] { "0+00".toU64Long } - intercept[NumberFormatException] { "+ffffffffffffffff".toU64Long } - } - - "not parse non-latin unicode digits" in { - // \u09e6 = BENGALI DIGIT ZERO (accepted by parseLong) - intercept[NumberFormatException] { "\u09e6".toU64Long } - } - - "parse mixed case" in { - assert(new RichU64String("aBcDeF").toU64Long == 0xabcdef) - } - - "not parse if negative" in { - val actual = intercept[NumberFormatException] { "-1".toU64Long } - assert(actual.getMessage == "For input string: \"-1\"") - - intercept[NumberFormatException] { "".toU64Long } - intercept[NumberFormatException] { "-f".toU64Long } - intercept[NumberFormatException] { "-af".toU64Long } - intercept[NumberFormatException] { "10-aff".toU64Long } - intercept[NumberFormatException] { "1-0aff".toU64Long } - } - - "not parse empty string" in { - // this is what Long.parseLong("") does - intercept[NumberFormatException] { "".toU64Long } - } - } -} diff --git a/util-core/src/test/scala/com/twitter/util/VarTest.scala b/util-core/src/test/scala/com/twitter/util/VarTest.scala index ef2a55b7cd..3bf1c3796a 100644 --- a/util-core/src/test/scala/com/twitter/util/VarTest.scala +++ b/util-core/src/test/scala/com/twitter/util/VarTest.scala @@ -1,15 +1,32 @@ package com.twitter.util import java.util.concurrent.atomic.AtomicReference -import org.junit.runner.RunWith import org.scalacheck.Arbitrary.arbitrary -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks import scala.collection.mutable +import org.scalatest.funsuite.AnyFunSuite + +class VarTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { + + // Workaround methods for dealing with Scala compiler warnings: + // + // assert(f.isDone) cannot be called directly because it results in a Scala compiler + // warning: 'possible missing interpolator: detected interpolated identifier `$conforms`' + // + // This is due to the implicit evidence required for `Future.isDone` which checks to see + // whether the Future that is attempting to have `isDone` called on it conforms to the type + // of `Future[Unit]`. This is done using `Predef.$conforms` + // https://www.scala-lang.org/api/2.12.2/scala/Predef$.html#$conforms[A]:A%3C:%3CA + // + // Passing that directly to `assert` causes problems because the `$conforms` is also seen as + // an interpolated string. We get around it by evaluating first and passing the result to + // `assert`. + private[this] def isDone(f: Future[Unit]): Boolean = + f.isDefined + + private[this] def assertIsDone(f: Future[Unit]): Unit = + assert(isDone(f)) -@RunWith(classOf[JUnitRunner]) -class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { private case class U[T](init: T) extends UpdatableVar[T](init) { import Var.Observer @@ -19,10 +36,12 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { override def observe(d: Int, obs: Observer[T]) = { accessCount += 1 observerCount += 1 - Closable.all(super.observe(d, obs), Closable.make { deadline => - observerCount -= 1 - Future.Done - }) + Closable.all( + super.observe(d, obs), + Closable.make { deadline => + observerCount -= 1 + Future.Done + }) } } @@ -34,9 +53,7 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(Var.sample(s) == "8923") var buf = mutable.Buffer[String]() - s.changes.register(Witness({ (v: String) => - buf += v; () - })) + s.changes.register(Witness({ (v: String) => buf += v; () })) assert(buf == Seq("8923")) v() = 111 assert(buf == Seq("8923", "111")) @@ -45,23 +62,13 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { test("depth ordering") { val v0 = U(3) val v1 = U(2) - val v2 = v1 flatMap { i => - v1 - } - val v3 = v2 flatMap { i => - v1 - } - val v4 = v3 flatMap { i => - v0 - } + val v2 = v1 flatMap { i => v1 } + val v3 = v2 flatMap { i => v1 } + val v4 = v3 flatMap { i => v0 } var result = 1 - v4.changes.register(Witness({ (i: Int) => - result = result + 2 - })) // result = 3 - v0.changes.register(Witness({ (i: Int) => - result = result * 2 - })) // result = 6 + v4.changes.register(Witness({ (i: Int) => result = result + 2 })) // result = 3 + v0.changes.register(Witness({ (i: Int) => result = result * 2 })) // result = 6 assert(result == 6) result = 1 // reset the value, but this time the ordering will go v0, v4 because of depth @@ -75,15 +82,15 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { val v1 = Var(2) var result = 0 - val o1 = v1.changes.register(Witness({ (i: Int) => - result = result + i - })) // result = 2 - val o2 = v1.changes.register(Witness({ (i: Int) => - result = result * i * i - })) // result = 2 * 2 * 2 = 8 - val o3 = v1.changes.register(Witness({ (i: Int) => - result = result + result + i - })) // result = 8 + 8 + 2 = 18 + val o1 = v1.changes.register(Witness({ (i: Int) => result = result + i })) // result = 2 + val o2 = + v1.changes.register(Witness({ (i: Int) => + result = result * i * i + })) // result = 2 * 2 * 2 = 8 + val o3 = + v1.changes.register(Witness({ (i: Int) => + result = result + result + i + })) // result = 8 + 8 + 2 = 18 assert(result == 18) // ensure those three things happened in sequence @@ -202,31 +209,27 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(vv == 111) assert(closed == Time.Zero) - val o1 = v.changes.register(Witness({ (v: Int) => - () - })) + val o1 = v.changes.register(Witness({ (v: Int) => () })) val t = Time.now val f = o.close(t) assert(called == 1) assert(closed == Time.Zero) - assert(f.isDone) + assertIsDone(f) // Closing the Var.async process is asynchronous with closing // the Var itself. val f1 = o1.close(t) assert(closed == t) - assert(f1.isDone) + assertIsDone(f1) } test("Var.collect: empty") { - assert(Var.collect(Seq.empty[Var[Int]]).sample == Seq.empty) + assert(Var.collect(List.empty[Var[Int]]).sample() == Nil) } test("Var.collect[Seq]") { - def ranged(n: Int) = Seq.tabulate(n) { i => - Var(i) - } + def ranged(n: Int) = Seq.tabulate(n) { i => Var(i) } for (i <- 1 to 10) { val vars = ranged(i) @@ -240,25 +243,6 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { } } - // This is either very neat or very horrendous, - // depending on your point of view. - test("Var.collect[Set]") { - val vars = Seq(Var(1), Var(2), Var(3)) - - val coll = Var.collect(vars.map { v => - v: Var[Int] - }.toSet) - val ref = new AtomicReference[Set[Int]] - coll.changes.register(Witness(ref)) - assert(ref.get == Set(1, 2, 3)) - - vars(1).update(1) - assert(ref.get == Set(1, 3)) - - vars(1).update(999) - assert(ref.get == Set(1, 999, 3)) - } - test("Var.collect: ordering") { val v1 = Var(1) val v2 = v1.map(_ * 2) @@ -273,20 +257,52 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { assert(ref.get == Seq((1, 2), (2, 4))) } - test("Var.collectIndependent: not overflowing the stack") { - val vars: Seq[Var[Int]] = for (i <- 1 to 10000) yield { - Var(i) // vars(i-1).map(_+2) - } + test("Var.collect overflows the stack") { + val vars: Seq[Var[Int]] = for (i <- 1 to 10 * 1000) yield Var(i) intercept[java.lang.StackOverflowError] { val v = Var.collect(vars) - v.observe { x => - } + v.observe { x => /* do nothing */ } } + } + test("Var.collectIndependent does NOT overflow the stack") { + val vars: Seq[Var[Int]] = for (i <- 1 to 10 * 1000) yield Var(i) val v = Var.collectIndependent(vars) - v.observe { x => - } + v.observe { x => /* do nothing */ } + assert(v.sample() == Seq.tabulate(10 * 1000)(i => i + 1)) + } + + test("Var.collectIndependent") { + val v1 = Var(10) + val v2 = Var(5) + val v3 = Var.collectIndependent(Seq[Var[Int]](v1, v2)) + + assert(v3.sample() == Seq(10, 5)) + + v1() = 11 + assert(v3.sample() == Seq(11, 5)) + + v2() = 6 + assert(v3.sample() == Seq(11, 6)) + } + + test("difference between collect and collectIndependent") { + val v1 = Var(1) + val v2 = v1.map(_ * 2) + val vCollect = Var.collect(Seq(v1, v2)).map { case Seq(a, b) => (a, b) } + val vCollectIndependent = Var.collectIndependent(Seq(v1, v2)).map { case Seq(a, b) => (a, b) } + + val refCollect = new AtomicReference[Seq[(Int, Int)]] + vCollect.changes.build.register(Witness(refCollect)) + + val refCollectIndependent = new AtomicReference[Seq[(Int, Int)]] + vCollectIndependent.changes.build.register(Witness(refCollectIndependent)) + + v1() = 2 + + assert(refCollect.get == Vector((1, 2), (2, 4))) + assert(refCollectIndependent.get == Vector((1, 2), (2, 2), (2, 4))) } /** @@ -296,9 +312,7 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { val contents = List(1, 2, 3, 4) val v1 = Var.value(contents) assert(Var.sample(v1) eq contents) - v1.changes.register(Witness({ (l: List[Int]) => - assert(contents eq l); () - })) + v1.changes.register(Witness({ (l: List[Int]) => assert(contents eq l); () })) } /** @@ -312,12 +326,8 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { Var.value(i * 2) } - val c1 = f.changes.register(Witness({ (i: Int) => - assert(i == 22); () - })) - val c2 = f.changes.register(Witness({ (i: Int) => - assert(i == 22); () - })) + val c1 = f.changes.register(Witness({ (i: Int) => assert(i == 22); () })) + val c2 = f.changes.register(Witness({ (i: Int) => assert(i == 22); () })) c1.close() c2.close() @@ -326,9 +336,7 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { v() = 22 // now it's safe to re-observe var observed = 3 - f.changes.register(Witness({ (i: Int) => - observed = i - })) + f.changes.register(Witness({ (i: Int) => observed = i })) assert(Var.sample(f) == 44) assert(Var.sample(v) == 22) @@ -347,12 +355,8 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { */ test("Var not executing until observed") { val x = Var(0) - val invertX = x map { i => - 1 / i - } - val result = x flatMap { i => - if (i == 0) Var(0) else invertX - } + val invertX = x map { i => 1 / i } + val result = x flatMap { i => if (i == 0) Var(0) else invertX } x() = 42 x() = 0 // this should not throw an exception because there are no observers @@ -381,13 +385,9 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { test("Don't propagate up-to-date " + typ + "-valued Var observations") { val v = Var(123) val w = newVar(333) - val x = v flatMap { _ => - w - } + val x = v flatMap { _ => w } var buf = mutable.Buffer[Int]() - x.changes.register(Witness({ (v: Int) => - buf += v; () - })) + x.changes.register(Witness({ (v: Int) => buf += v; () })) assert(buf == Seq(333)) v() = 333 @@ -404,9 +404,7 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { } var buf = mutable.Buffer[Int]() - x.changes.register(Witness({ (v: Int) => - buf += v; () - })) + x.changes.register(Witness({ (v: Int) => buf += v; () })) assert(buf == Seq(333)) v() = 333 @@ -470,7 +468,8 @@ class VarTest extends FunSuite with GeneratorDrivenPropertyChecks { // in order to avoid extremely long runtimes, we limit the max size // of the generated Seqs and Sets. A limit of 50 keeps runtime below // 1 second on my crufty old laptop. - implicit val generatorDrivenConfig = PropertyCheckConfiguration(sizeRange = 50) + implicit val generatorDrivenConfig: PropertyCheckConfiguration = + PropertyCheckConfiguration(sizeRange = 50) forAll(arbitrary[Seq[Set[Int]]].suchThat(_.nonEmpty)) { sets => val v = Var(sets.head) diff --git a/util-core/src/test/scala/com/twitter/util/WaitQueueTest.scala b/util-core/src/test/scala/com/twitter/util/WaitQueueTest.scala index 4e96878bdc..f4b5a7474d 100644 --- a/util-core/src/test/scala/com/twitter/util/WaitQueueTest.scala +++ b/util-core/src/test/scala/com/twitter/util/WaitQueueTest.scala @@ -1,12 +1,12 @@ package com.twitter.util import org.scalacheck.Gen -import org.scalatest.FunSuite -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -class WaitQueueTest extends FunSuite with GeneratorDrivenPropertyChecks { +class WaitQueueTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { - private class EmptyK[A] extends Promise.K[A] { + private class EmptyK[A] extends Promise.K[A](Local.Context.empty) { protected[util] def depth: Short = 0 def apply(ta: Try[A]): Unit = () } @@ -21,9 +21,7 @@ class WaitQueueTest extends FunSuite with GeneratorDrivenPropertyChecks { } yield genWaitQueueOf(size) test("size") { - forAll(Gen.choose(0, 100)) { size => - assert(genWaitQueueOf(size).size == size) - } + forAll(Gen.choose(0, 100)) { size => assert(genWaitQueueOf(size).size == size) } } test("contains") { diff --git a/util-core/src/test/scala/com/twitter/util/WindowedAdderTest.scala b/util-core/src/test/scala/com/twitter/util/WindowedAdderTest.scala index f057afba01..c397ca5865 100644 --- a/util-core/src/test/scala/com/twitter/util/WindowedAdderTest.scala +++ b/util-core/src/test/scala/com/twitter/util/WindowedAdderTest.scala @@ -1,12 +1,9 @@ package com.twitter.util -import com.twitter.conversions.time._ -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import com.twitter.conversions.DurationOps._ +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class WindowedAdderTest extends FunSuite { +class WindowedAdderTest extends AnyFunSuite { private def newAdder() = WindowedAdder(3 * 1000, 3, Stopwatch.timeMillis) test("sums things up when time stands still") { diff --git a/util-core/src/test/scala/com/twitter/util/WorkQueueFiberTest.scala b/util-core/src/test/scala/com/twitter/util/WorkQueueFiberTest.scala new file mode 100644 index 0000000000..537ee96f0f --- /dev/null +++ b/util-core/src/test/scala/com/twitter/util/WorkQueueFiberTest.scala @@ -0,0 +1,450 @@ +package com.twitter.util + +import com.twitter.conversions.DurationOps.richDurationFromInt +import java.util.concurrent.CountDownLatch +import java.util.concurrent.atomic.AtomicInteger +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.mutable.ListBuffer + +class WorkQueueFiberTest extends AnyFunSuite { + + private def await[T](t: Awaitable[T]): T = Await.result(t, 10.seconds) + + private class TestMetrics extends WorkQueueFiber.FiberMetrics { + + val fiberFlushed, fibersCreated, threadLocalSubmits, taskSubmissions, schedulerSubmissions = + new AtomicInteger() + + def flushIncrement(): Unit = fiberFlushed.incrementAndGet() + + def fiberCreated(): Unit = fibersCreated.incrementAndGet() + def threadLocalSubmitIncrement(): Unit = threadLocalSubmits.incrementAndGet() + def taskSubmissionIncrement(): Unit = taskSubmissions.incrementAndGet() + def schedulerSubmissionIncrement(): Unit = schedulerSubmissions.incrementAndGet() + } + + private class TestFiber(metrics: TestMetrics) extends WorkQueueFiber(metrics) { + + var runnable: Runnable = null + + def executePendingRunnable(): Unit = { + assert(runnable != null) + val r = runnable + runnable = null + r.run() + } + + protected def schedulerSubmit(r: Runnable): Unit = { + assert(runnable == null) + runnable = r + } + + override protected def schedulerFlush(): Unit = () + } + + test("fiber creations are tracked") { + val N = 5 + val metrics = new TestMetrics + for (_ <- 0 until N) { + new TestFiber(metrics) + } + assert(metrics.fibersCreated.get == N) + } + + test("submit multiple tasks while fiber not running") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + assert(metrics.fibersCreated.get == 1) + + val executed = new AtomicInteger() + + val N = 5 + for (_ <- 0 until N) { + f.submitTask { () => + executed.incrementAndGet() + } + } + + assert(f.runnable != null) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.taskSubmissions.get == N) + assert(metrics.schedulerSubmissions.get == 1) + assert(executed.get == 0) + + f.executePendingRunnable() + assert(executed.get == N) + } + + test("submit a thread local task") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + assert(metrics.fibersCreated.get == 1) + + val executed = new AtomicInteger() + val submissionOrder = new ListBuffer[Int] + + f.submitTask { () => + executed.incrementAndGet() + + // This one should be submitted thread-locally. + f.submitTask { () => + executed.incrementAndGet() + submissionOrder += 1 + } + + // fiber task 1 should complete all the way through before we execute the sub-task + submissionOrder += 0 + } + + assert(f.runnable != null) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + assert(executed.get == 0) + + f.executePendingRunnable() + assert(metrics.threadLocalSubmits.get == 1) + assert(metrics.taskSubmissions.get == 2) + assert(executed.get == 2) + assert(submissionOrder.result() == List(0, 1)) + + // Now throw in one more to make sure we don't get a thread-local submit since the fiber + // is clearly not executing. + f.submitTask { () => + executed.incrementAndGet() + submissionOrder += 2 + } + + assert(metrics.threadLocalSubmits.get == 1) + assert(metrics.taskSubmissions.get == 3) + assert(metrics.schedulerSubmissions.get == 2) + f.executePendingRunnable() + + assert(executed.get == 3) + assert(submissionOrder.result() == List(0, 1, 2)) + } + + test("submit a task while a separate thread is executing the fiber") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + assert(metrics.fibersCreated.get == 1) + + val taskAwait, taskExecuting = new CountDownLatch(1) + + f.submitTask { () => + taskExecuting.countDown() + taskAwait.await() + } + + assert(f.runnable != null) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + + // Start the fiber in a separate thread + val t = new Thread({ () => + f.executePendingRunnable() + }) + t.start() + + taskExecuting.await() + + val executingThread = new Promise[Thread]() + f.submitTask { () => + executingThread.setValue(Thread.currentThread) + } + + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.taskSubmissions.get == 2) + assert(metrics.threadLocalSubmits.get == 0) + assert(!executingThread.isDefined) + + taskAwait.countDown() + assert(await(executingThread) == t) + } + + test("flushing re-submits to the scheduler") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + val latch1, latch2, latch3, flushed = new CountDownLatch(1) + + // This will eventually be the test thread so we don't need it to be volatile. + var threadLocalSubmissionsExecutionThread: Thread = null + var lastTaskDone = false + f.submitTask { () => + latch1.countDown() + latch2.await() + f.submitTask { () => + threadLocalSubmissionsExecutionThread = Thread.currentThread() + } + + // Now flush and continue to suspend the thread. + assert(metrics.fiberFlushed.get == 0) + f.flush() + assert(metrics.fiberFlushed.get == 1) + flushed.countDown() + latch3.await() + + // schedule one more task that should not be a thread local submit because flush + // voided that. + f.submitTask { () => + lastTaskDone = true + } + } + + assert(metrics.taskSubmissions.get == 1) + + val t = new Thread({ () => + f.executePendingRunnable() + }) + t.start() + + latch1.await() + + // We'll run this one in the test thread but it's going to be submitted as a non-thread local + var done = false + f.submitTask { () => + done = true + } + + // Release thread t and then stall t right after the flush + latch2.countDown() + flushed.await() + + assert(metrics.taskSubmissions.get == 3) + assert(metrics.threadLocalSubmits.get == 1) + assert(metrics.schedulerSubmissions.get == 2) + assert(threadLocalSubmissionsExecutionThread == null) + assert(!done) + + // We should now have 2 tasks awaiting whoever executes the next submission. Now run it. + f.executePendingRunnable() + assert(threadLocalSubmissionsExecutionThread == Thread.currentThread()) + assert(done) + + // Now let our thread submit it's final task which shouldn't be a thread local task even + // though it's submitted from within a FiberTask. + latch3.countDown() + t.join() + + assert(metrics.taskSubmissions.get == 4) + assert(metrics.threadLocalSubmits.get == 1) + assert(metrics.schedulerSubmissions.get == 3) + assert(!lastTaskDone) + + f.executePendingRunnable() + assert(lastTaskDone) + } + + test("flushing to a recursive scheduler") { + val metrics = new TestMetrics + val f = new WorkQueueFiber(metrics) { + override protected def schedulerSubmit(r: Runnable): Unit = { + r.run() + } + override protected def schedulerFlush(): Unit = () + } + + f.submitTask { () => + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.threadLocalSubmits.get == 0) + + var threadLocalSubmit = false + f.submitTask { () => + threadLocalSubmit = true + } + assert(metrics.taskSubmissions.get == 2) + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.threadLocalSubmits.get == 1) + + assert(!threadLocalSubmit) + assert(metrics.fiberFlushed.get == 0) + f.flush() + assert(metrics.fiberFlushed.get == 1) + assert(threadLocalSubmit) + assert(metrics.taskSubmissions.get == 2) + assert(metrics.schedulerSubmissions.get == 2) + assert(metrics.threadLocalSubmits.get == 1) + + // After the flush() call, even recursive calls should act like an external submission. + var shouldExecuteNow = false + f.submitTask { () => + shouldExecuteNow = true + } + assert(shouldExecuteNow) + assert(metrics.taskSubmissions.get == 3) + assert(metrics.schedulerSubmissions.get == 3) + assert(metrics.threadLocalSubmits.get == 1) + } + } + + test("thread-local flush with empty work queues still de-schedules itself") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + var executed = false + f.submitTask { () => + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.fiberFlushed.get == 0) + + f.flush() + + // At this point we should be as if we're not actually executing in the fiber. + + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.fiberFlushed.get == 1) + + // This will get added to the end of the queue and we'll run it without + // needing a resubmit. + f.submitTask { () => + executed = true + } + + assert(!executed) + assert(metrics.taskSubmissions.get == 2) + assert(metrics.schedulerSubmissions.get == 2) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.fiberFlushed.get == 1) + } + + f.executePendingRunnable() + // Still not executed. Even though the submission was done inside a FiberTask + // the flush() had caused the task to no longer see itself as the executor of + // the fiber. + assert(!executed) + + f.executePendingRunnable() + assert(executed) + } + + test("a flush outside the executing code is a no-op for fiber") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + var executed = false + f.submitTask { () => + executed = true + } + + assert(metrics.taskSubmissions.get == 1) + assert(metrics.schedulerSubmissions.get == 1) + assert(metrics.threadLocalSubmits.get == 0) + assert(metrics.fiberFlushed.get == 0) + + f.flush() + assert(metrics.fiberFlushed.get == 0) + assert(!executed) + + f.executePendingRunnable() + assert(executed) + } + + test("a flush triggers a scheduler flush whether or not its in the executing code") { + val metrics = new TestMetrics + var schedulerFlushes = 0 + val f = new WorkQueueFiber(metrics) { + override protected def schedulerSubmit(r: Runnable): Unit = { + r.run() + } + override protected def schedulerFlush(): Unit = { + schedulerFlushes += 1 + } + } + // flush from within executing code + f.submitTask(() => { + f.flush() + }) + assert(schedulerFlushes == 1) + assert(metrics.fiberFlushed.get == 1) + // flush from outside executing code + f.flush() + assert(schedulerFlushes == 2) + // the fiber didn't flush this time + assert(metrics.fiberFlushed.get == 1) + } + + test("fiber tracks resources") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + val executed = new AtomicInteger() + val N = 5 + for (_ <- 0 until N) { + f.submitTask { () => + executed.incrementAndGet() + } + } + + // The fiber hasn't done any work yet + assert(executed.get == 0) + assert(f.tracker.totalCpuTime() == 0) + assert(f.tracker.numContinuations() == 0) + + f.executePendingRunnable() + assert(executed.get == N) + // The resource tracker records the work done by the fiber + assert(f.tracker.totalCpuTime() > 0) + assert(f.tracker.numContinuations() == N) + } + + test("resource tracking with flush") { + val metrics = new TestMetrics + val f = new TestFiber(metrics) + + val latch1, latch2, flushed = new CountDownLatch(1) + + var innerTaskDone = false + f.submitTask { () => + latch1.countDown() + latch2.await() + f.submitTask { () => + innerTaskDone = true + } + + // Now flush + f.flush() + flushed.countDown() + } + + // Nothing recorded yet by the resource tracker + assert(f.tracker.totalCpuTime() == 0) + assert(f.tracker.numContinuations() == 0) + + val t = new Thread({ () => + f.executePendingRunnable() + }) + t.start() + + latch1.await() + // At this point nothing has been recorded with the resource tracker yet + assert(f.tracker.totalCpuTime() == 0) + assert(f.tracker.numContinuations() == 0) + + // Release thread t and then stall t right after the flush + latch2.countDown() + flushed.await() + t.join() + // Now the cpu time and the number of continuations should have been updated + val cpuTimeAfterFLush = f.tracker.totalCpuTime() + assert(cpuTimeAfterFLush > 0) + assert(f.tracker.numContinuations() == 1) + + assert(!innerTaskDone) + // We should now have the inner task awaiting whoever executes the next submission. Now run it. + f.executePendingRunnable() + assert(innerTaskDone) + // More cpu time should have been added as well as one more continuation + val cpuTimeAfterInnerTask = f.tracker.totalCpuTime() + assert(cpuTimeAfterInnerTask > cpuTimeAfterFLush) + assert(f.tracker.numContinuations() == 2) + } +} diff --git a/util-function/BUILD b/util-function/BUILD index 82d3156556..d316723ed2 100644 --- a/util-function/BUILD +++ b/util-function/BUILD @@ -1,4 +1,5 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-function/src/main/java", ], diff --git a/util-function/src/main/java/BUILD b/util-function/src/main/java/BUILD index 4e3a84dba8..969c923223 100644 --- a/util-function/src/main/java/BUILD +++ b/util-function/src/main/java/BUILD @@ -1,9 +1,6 @@ -java_library( - sources = rglobs("*.java"), - fatal_warnings = True, - provides = artifact( - org = "com.twitter", - name = "util-function-java", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-function/src/main/java/com/twitter/function", + ], ) diff --git a/util-function/src/main/java/com/twitter/function/BUILD b/util-function/src/main/java/com/twitter/function/BUILD new file mode 100644 index 0000000000..ce03a4870f --- /dev/null +++ b/util-function/src/main/java/com/twitter/function/BUILD @@ -0,0 +1,11 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-function-java", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-hashing/BUILD b/util-hashing/BUILD index cc89f2ca64..62d56876e1 100644 --- a/util-hashing/BUILD +++ b/util-hashing/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-hashing/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-hashing/src/test/scala", ], diff --git a/util-hashing/src/main/scala/BUILD b/util-hashing/src/main/scala/BUILD index 0a048c193f..e55f18e86a 100644 --- a/util-hashing/src/main/scala/BUILD +++ b/util-hashing/src/main/scala/BUILD @@ -1,9 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-hashing", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-hashing/src/main/scala/com/twitter/hashing", + ], ) diff --git a/util-hashing/src/main/scala/com/twitter/hashing/BUILD b/util-hashing/src/main/scala/com/twitter/hashing/BUILD new file mode 100644 index 0000000000..b4531eb750 --- /dev/null +++ b/util-hashing/src/main/scala/com/twitter/hashing/BUILD @@ -0,0 +1,11 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-hashing", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-hashing/src/main/scala/com/twitter/hashing/KetamaDistributor.scala b/util-hashing/src/main/scala/com/twitter/hashing/ConsistentHashingDistributor.scala similarity index 61% rename from util-hashing/src/main/scala/com/twitter/hashing/KetamaDistributor.scala rename to util-hashing/src/main/scala/com/twitter/hashing/ConsistentHashingDistributor.scala index ba12f5e73e..5fdb7503b2 100644 --- a/util-hashing/src/main/scala/com/twitter/hashing/KetamaDistributor.scala +++ b/util-hashing/src/main/scala/com/twitter/hashing/ConsistentHashingDistributor.scala @@ -1,8 +1,9 @@ package com.twitter.hashing import java.security.MessageDigest +import java.util -case class KetamaNode[A](identifier: String, weight: Int, handle: A) +case class HashNode[A](identifier: String, weight: Int, handle: A) /** * @note Certain versions of libmemcached return subtly different results for points on @@ -10,11 +11,11 @@ case class KetamaNode[A](identifier: String, weight: Int, handle: A) * clients who depend on those versions of libmemcached, we have to reproduce their result. * If the `oldLibMemcachedVersionComplianceMode` is `true` the behavior will be reproduced. */ -class KetamaDistributor[A]( - ketamaNodes: Seq[KetamaNode[A]], +class ConsistentHashingDistributor[A]( + hashNodes: Seq[HashNode[A]], numReps: Int, - oldLibMemcachedVersionComplianceMode: Boolean = false -) extends Distributor[A] { + oldLibMemcachedVersionComplianceMode: Boolean = false) + extends Distributor[A] { /** * Extract a little-endian integer from a given byte array starting at `offset`. @@ -53,15 +54,15 @@ class KetamaDistributor[A]( md.update(('0' + j).toByte) } - private[this] val continuum: java.util.TreeMap[Long, KetamaNode[A]] = { + private[this] val (continuumKeys: Array[Long], continuumValues: Array[HashNode[A]]) = { val md5 = MessageDigest.getInstance("MD5") - val underlying = new java.util.TreeMap[Long, KetamaNode[A]]() + val underlying = new java.util.TreeMap[Long, HashNode[A]]() val dash: Byte = '-'.toByte - val nodeCount = ketamaNodes.size - val totalWeight = ketamaNodes.foldLeft(0) { _ + _.weight } + val nodeCount = hashNodes.size + val totalWeight = hashNodes.foldLeft(0) { _ + _.weight } - val it = ketamaNodes.iterator + val it = hashNodes.iterator while (it.hasNext) { val node = it.next() @@ -98,29 +99,51 @@ class KetamaDistributor[A]( assert(underlying.size >= numReps * (nodeCount - 1)) } - underlying - } + val keys = new Array[Long](underlying.size()) + val values = new Array[HashNode[A]](underlying.size()) - def nodes: Seq[A] = ketamaNodes.map(_.handle) - def nodeCount: Int = ketamaNodes.size + var i = 0; + val underlyingIterator = underlying.entrySet().iterator() + while (underlyingIterator.hasNext) { + val entry = underlyingIterator.next() + keys(i) = entry.getKey + values(i) = entry.getValue + i += 1 + } - private def mapEntryForHash(hash: Long): java.util.Map.Entry[Long, KetamaNode[A]] = { - // hashes are 32-bit because they are 32-bit on the libmemcached and - // we need to maintain compatibility with libmemcached - val truncatedHash = hash & 0xffffffffL + (keys, values) + } - val entry = continuum.ceilingEntry(truncatedHash) - if (entry == null) - continuum.firstEntry - else - entry + def nodes: Seq[A] = hashNodes.map(_.handle) + def nodeCount: Int = hashNodes.size + + // hashes are 32-bit because they are 32-bit on the libmemcached and + // we need to maintain compatibility with libmemcached + private[this] def truncateHash(hash: Long): Long = hash & 0xffffffffL + + private def mapEntryForHash(hash: Long): Int = { + val truncatedHash = truncateHash(hash) + val index = util.Arrays.binarySearch(continuumKeys, truncatedHash) + if (index >= 0) { + index + } else { + val insertionPoint = -(index + 1) + if (insertionPoint < continuumValues.length) { + insertionPoint + } else { + 0 + } + } } + def partitionIdForHash(hash: Long): Long = + continuumKeys(mapEntryForHash(hash)) + def entryForHash(hash: Long): (Long, A) = { - val entry = mapEntryForHash(hash) - (entry.getKey, entry.getValue.handle) + val index = mapEntryForHash(hash) + (continuumKeys(index), continuumValues(index).handle) } def nodeForHash(hash: Long): A = - mapEntryForHash(hash).getValue.handle + continuumValues(mapEntryForHash(hash)).handle } diff --git a/util-hashing/src/main/scala/com/twitter/hashing/Distributor.scala b/util-hashing/src/main/scala/com/twitter/hashing/Distributor.scala index c90e0774d2..cc8f6d0bdc 100644 --- a/util-hashing/src/main/scala/com/twitter/hashing/Distributor.scala +++ b/util-hashing/src/main/scala/com/twitter/hashing/Distributor.scala @@ -2,14 +2,16 @@ package com.twitter.hashing trait Distributor[A] { def entryForHash(hash: Long): (Long, A) + def partitionIdForHash(hash: Long): Long def nodeForHash(hash: Long): A def nodeCount: Int def nodes: Seq[A] } class SingletonDistributor[A](node: A) extends Distributor[A] { - def entryForHash(hash: Long) = (hash, node) - def nodeForHash(hash: Long) = node + def entryForHash(hash: Long): (Long, A) = (hash, node) + def partitionIdForHash(hash: Long): Long = hash + def nodeForHash(hash: Long): A = node def nodeCount = 1 - def nodes = Seq(node) + def nodes: Seq[A] = Seq(node) } diff --git a/util-hashing/src/main/scala/com/twitter/hashing/Hashable.scala b/util-hashing/src/main/scala/com/twitter/hashing/Hashable.scala index 24e6184618..9ffc61f392 100644 --- a/util-hashing/src/main/scala/com/twitter/hashing/Hashable.scala +++ b/util-hashing/src/main/scala/com/twitter/hashing/Hashable.scala @@ -17,15 +17,11 @@ trait Hashable[-T, +R] extends (T => R) { self => trait LowPriorityHashable { // XOR the high 32 bits into the low to get a int: implicit def toInt[T](implicit h: Hashable[T, Long]): Hashable[T, Int] = - h.andThen { long => - (long >> 32).toInt ^ long.toInt - } + h.andThen { long => (long >> 32).toInt ^ long.toInt } // Get the UTF-8 bytes of a string to hash it implicit def fromString[T](implicit h: Hashable[Array[Byte], T]): Hashable[String, T] = - h.compose { s: String => - s.getBytes - } + h.compose { (s: String) => s.getBytes } } object Hashable extends LowPriorityHashable { @@ -43,11 +39,11 @@ object Hashable extends LowPriorityHashable { // Some standard hashing: def hashCode[T]: Hashable[T, Int] = new Hashable[T, Int] { def apply(t: T) = t.hashCode } - private[this] val MaxUnsignedInt: Long = 0xFFFFFFFFL + private[this] val MaxUnsignedInt: Long = 0xffffffffL /** * FNV fast hashing algorithm in 32 bits. - * @see http://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash + * @see https://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash */ val FNV1_32: Hashable[Array[Byte], Int] = new Hashable[Array[Byte], Int] { def apply(key: Array[Byte]) = { @@ -67,7 +63,7 @@ object Hashable extends LowPriorityHashable { /** * FNV fast hashing algorithm in 32 bits, variant with operations reversed. - * @see http://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash + * @see https://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash */ val FNV1A_32: Hashable[Array[Byte], Int] = new Hashable[Array[Byte], Int] { def apply(key: Array[Byte]): Int = { @@ -87,7 +83,7 @@ object Hashable extends LowPriorityHashable { /** * FNV fast hashing algorithm in 64 bits. - * @see http://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash + * @see https://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash */ val FNV1_64: Hashable[Array[Byte], Long] = new Hashable[Array[Byte], Long] { def apply(key: Array[Byte]): Long = { @@ -107,7 +103,7 @@ object Hashable extends LowPriorityHashable { /** * FNV fast hashing algorithm in 64 bits, variant with operations reversed. - * @see http://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash + * @see https://en.wikipedia.org/wiki/Fowler_Noll_Vo_hash */ val FNV1A_64: Hashable[Array[Byte], Long] = new Hashable[Array[Byte], Long] { def apply(key: Array[Byte]): Long = { @@ -179,7 +175,7 @@ object Hashable extends LowPriorityHashable { /** * Paul Hsieh's hash function. - * http://www.azillionmonkeys.com/qed/hash.html + * https://www.azillionmonkeys.com/qed/hash.html */ val HSIEH: Hashable[Array[Byte], Int] = new Hashable[Array[Byte], Int] { def apply(key: Array[Byte]): Int = { @@ -244,7 +240,7 @@ object Hashable extends LowPriorityHashable { /** * Jenkins Hash Function - * http://en.wikipedia.org/wiki/Jenkins_hash_function + * https://en.wikipedia.org/wiki/Jenkins_hash_function */ val JENKINS: Hashable[Array[Byte], Long] = new Hashable[Array[Byte], Long] { def apply(key: Array[Byte]): Long = { diff --git a/util-hashing/src/test/resources/BUILD b/util-hashing/src/test/resources/BUILD index e0afa86adb..07015f79b6 100644 --- a/util-hashing/src/test/resources/BUILD +++ b/util-hashing/src/test/resources/BUILD @@ -1,6 +1,8 @@ resources( - sources = rglobs( - "*", - exclude = [rglobs("BUILD*", "*.pyc")], - ), + sources = [ + "!**/*.pyc", + "!BUILD*", + "**/*", + ], + tags = ["bazel-compatible"], ) diff --git a/util-hashing/src/test/scala/BUILD b/util-hashing/src/test/scala/BUILD index 80c1ef9ac4..403d5d45a0 100644 --- a/util-hashing/src/test/scala/BUILD +++ b/util-hashing/src/test/scala/BUILD @@ -1,13 +1,6 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalacheck", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-hashing/src/main/scala", - "util/util-hashing/src/test/resources", + "util/util-hashing/src/test/scala/com/twitter/hashing", ], ) diff --git a/util-hashing/src/test/scala/com/twitter/hashing/BUILD b/util-hashing/src/test/scala/com/twitter/hashing/BUILD new file mode 100644 index 0000000000..c0b43f401e --- /dev/null +++ b/util-hashing/src/test/scala/com/twitter/hashing/BUILD @@ -0,0 +1,20 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-hashing/src/main/scala/com/twitter/hashing", + "util/util-hashing/src/test/resources", + ], +) diff --git a/util-hashing/src/test/scala/com/twitter/hashing/ConsistentHashingDistributorTest.scala b/util-hashing/src/test/scala/com/twitter/hashing/ConsistentHashingDistributorTest.scala new file mode 100644 index 0000000000..8a2f2bc833 --- /dev/null +++ b/util-hashing/src/test/scala/com/twitter/hashing/ConsistentHashingDistributorTest.scala @@ -0,0 +1,120 @@ +package com.twitter.hashing + +import org.scalacheck.Gen +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import scala.collection.mutable +import java.io.BufferedReader +import java.io.InputStreamReader +import java.nio.ByteBuffer +import java.nio.ByteOrder +import java.security.MessageDigest +import org.scalatest.funsuite.AnyFunSuite + +class ConsistentHashingDistributorTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { + val nodes = Seq( + HashNode("10.0.1.1", 600, 1), + HashNode("10.0.1.2", 300, 2), + HashNode("10.0.1.3", 200, 3), + HashNode("10.0.1.4", 350, 4), + HashNode("10.0.1.5", 1000, 5), + HashNode("10.0.1.6", 800, 6), + HashNode("10.0.1.7", 950, 7), + HashNode("10.0.1.8", 100, 8) + ) + + // 160 is the hard coded value for libmemcached, which was this input data is from + val ketamaDistributor = new ConsistentHashingDistributor(nodes, 160) + val ketamaDistributorInoldLibMemcachedVersionComplianceMode = + new ConsistentHashingDistributor(nodes, 160, true) + + test("KetamaDistributor should pick the correct node with ketama hash function") { + // Test from Smile's KetamaNodeLocatorSpec.scala + + // Load known good results (key, hash(?), continuum ceiling(?), IP) + val stream = getClass.getClassLoader.getResourceAsStream("ketama_results") + val reader = new BufferedReader(new InputStreamReader(stream)) + val expected = new mutable.ListBuffer[Array[String]] + var line: String = reader.readLine() + while (line != null) { + val segments = line.split(" ") + assert(segments.length == 4) + expected += segments + line = reader.readLine() + } + assert(expected.size == 99) + + // Test that ketamaClient.clientOf(key) == expected IP + val handleToIp = nodes.map { n => n.handle -> n.identifier }.toMap + for (testcase <- expected) { + val hash = KeyHasher.KETAMA.hashKey(testcase(0).getBytes) + + val handle = ketamaDistributor.nodeForHash(hash) + val handle2 = ketamaDistributorInoldLibMemcachedVersionComplianceMode.nodeForHash(hash) + val resultIp = handleToIp(handle) + val resultIp2 = handleToIp(handle2) + assert(testcase(3) == resultIp) + assert(testcase(3) == resultIp2) + } + } + + test("KetamaDistributor should pick the correct node with 64-bit hash values") { + val knownGoodValues = Map( + -166124121512512L -> 5, + 8796093022208L -> 3, + 4312515125124L -> 2, + -8192481414141L -> 1, + -9515121512312L -> 5 + ) + + knownGoodValues foreach { + case (key, node) => + val handle = ketamaDistributor.nodeForHash(key) + assert(handle == node) + } + } + + test("KetamaDistributor should hashInt") { + val ketama = new ConsistentHashingDistributor[Unit](Seq.empty, 0, false) + forAll(Gen.chooseNum(0, Int.MaxValue)) { i => + val md = MessageDigest.getInstance("MD5") + ketama.hashInt(i, md) + val array = md.digest() + assert(array.toSeq == md.digest(i.toString.getBytes("UTF-8")).toSeq) + } + } + + test("KetamaDistributor should byteArrayToLE") { + val ketama = new ConsistentHashingDistributor[Unit](Seq.empty, 0, false) + forAll { (s: String) => + val ba = MessageDigest.getInstance("MD5").digest(s.getBytes("UTF-8")) + val bb = ByteBuffer.wrap(ba).order(ByteOrder.LITTLE_ENDIAN) + + assert(ketama.byteArrayToLE(ba, 0) == bb.getInt(0)) + assert(ketama.byteArrayToLE(ba, 4) == bb.getInt(4)) + assert(ketama.byteArrayToLE(ba, 8) == bb.getInt(8)) + assert(ketama.byteArrayToLE(ba, 12) == bb.getInt(12)) + } + } + + test( + "KetamaDistributor should partitionIdForHash is the same as entryForHash 1st tuple element") { + val keyHasher = KeyHasher.MURMUR3 + + { + // Special case this unicode string since this test has failed + // on it in the past. + val s = "ퟲ狂✟풻棇埲蟂덧➀缘็滀佳黟韕숻᩿볯箿䜐㫌홻ᖛ磌ᡎ油" + val hash = keyHasher.hashKey(s.getBytes("UTF-8")) + val (pid1, _) = ketamaDistributor.entryForHash(hash) + val pid2 = ketamaDistributor.partitionIdForHash(hash) + assert(pid1 == pid2) + } + + forAll { (s: String) => + val hash = keyHasher.hashKey(s.getBytes("UTF-8")) + val (pid1, _) = ketamaDistributor.entryForHash(hash) + val pid2 = ketamaDistributor.partitionIdForHash(hash) + assert(pid1 == pid2) + } + } +} diff --git a/util-hashing/src/test/scala/com/twitter/hashing/DiagnosticsTest.scala b/util-hashing/src/test/scala/com/twitter/hashing/DiagnosticsTest.scala index 502daea63d..8360613ab2 100644 --- a/util-hashing/src/test/scala/com/twitter/hashing/DiagnosticsTest.scala +++ b/util-hashing/src/test/scala/com/twitter/hashing/DiagnosticsTest.scala @@ -1,42 +1,37 @@ package com.twitter.hashing -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class DiagnosticsTest extends WordSpec { - "Diagnostics" should { - "print distribution" in { - val hosts = 1 until 500 map { "10.1.1." + _ + ":11211:4" } +class DiagnosticsTest extends AnyFunSuite { + test("Diagnostics should print distribution") { + val hosts = 1 until 500 map { "10.1.1." + _ + ":11211:4" } - val nodes = hosts.map { s => - val Array(host, port, weight) = s.split(":") - val identifier = host + ":" + port - KetamaNode(identifier, weight.toInt, identifier) - } + val nodes = hosts.map { s => + val Array(host, port, weight) = s.split(":") + val identifier = host + ":" + port + HashNode(identifier, weight.toInt, identifier) + } - val hashFunctions = List( - "FNV1_32" -> KeyHasher.FNV1_32, - "FNV1A_32" -> KeyHasher.FNV1A_32, - "FNV1_64" -> KeyHasher.FNV1_64, - "FNV1A_64" -> KeyHasher.FNV1A_64, - "CRC32-ITU" -> KeyHasher.CRC32_ITU, - "HSIEH" -> KeyHasher.HSIEH, - "JENKINS" -> KeyHasher.JENKINS, - "MURMUR3" -> KeyHasher.MURMUR3 - ) + val hashFunctions = List( + "FNV1_32" -> KeyHasher.FNV1_32, + "FNV1A_32" -> KeyHasher.FNV1A_32, + "FNV1_64" -> KeyHasher.FNV1_64, + "FNV1A_64" -> KeyHasher.FNV1A_64, + "CRC32-ITU" -> KeyHasher.CRC32_ITU, + "HSIEH" -> KeyHasher.HSIEH, + "JENKINS" -> KeyHasher.JENKINS, + "MURMUR3" -> KeyHasher.MURMUR3 + ) - val keys = (1 until 1000000).map(_.toString).toList - // Comment this out unless you're running it results - // hashFunctions foreach { case (s, h) => - // val distributor = new KetamaDistributor(nodes, 160) - // val tester = new DistributionTester(distributor) - // val start = Time.now - // val dev = tester.distributionDeviation(keys map { s => h.hashKey(s.getBytes) }) - // val duration = (Time.now - start).inMilliseconds - // println("%s\n distribution: %.5f\n duration: %dms\n".format(s, dev, duration)) - // } - } + val keys = (1 until 1000000).map(_.toString).toList + // Comment this out unless you're running it results + // hashFunctions foreach { case (s, h) => + // val distributor = new KetamaDistributor(nodes, 160) + // val tester = new DistributionTester(distributor) + // val start = Time.now + // val dev = tester.distributionDeviation(keys map { s => h.hashKey(s.getBytes) }) + // val duration = (Time.now - start).inMilliseconds + // println("%s\n distribution: %.5f\n duration: %dms\n".format(s, dev, duration)) + // } } } diff --git a/util-hashing/src/test/scala/com/twitter/hashing/DistributionTester.scala b/util-hashing/src/test/scala/com/twitter/hashing/DistributionTester.scala index 5688d70c8f..3f0ee37b58 100644 --- a/util-hashing/src/test/scala/com/twitter/hashing/DistributionTester.scala +++ b/util-hashing/src/test/scala/com/twitter/hashing/DistributionTester.scala @@ -15,13 +15,9 @@ class DistributionTester[A](distributor: Distributor[A]) { keysPerNode(key) += 1 } var frequencies = keysPerNode.values.toList - frequencies ++= 0 until (distributor.nodeCount - frequencies.size) map { _ => - 0 - } + frequencies ++= 0 until (distributor.nodeCount - frequencies.size) map { _ => 0 } val average = frequencies.sum.toDouble / frequencies.size - val diffs = frequencies.map { v => - math.pow((v - average), 2) - } + val diffs = frequencies.map { v => math.pow((v - average), 2) } val sd = math.sqrt(diffs.sum / (frequencies.size - 1)) sd / average } diff --git a/util-hashing/src/test/scala/com/twitter/hashing/HashableTest.scala b/util-hashing/src/test/scala/com/twitter/hashing/HashableTest.scala index 43c8715111..2a76d0196f 100644 --- a/util-hashing/src/test/scala/com/twitter/hashing/HashableTest.scala +++ b/util-hashing/src/test/scala/com/twitter/hashing/HashableTest.scala @@ -1,12 +1,9 @@ package com.twitter.hashing -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.FunSuite -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class HashableTest extends FunSuite with GeneratorDrivenPropertyChecks { +class HashableTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { private[this] val algorithms = Seq( Hashable.CRC32_ITU, @@ -21,9 +18,7 @@ class HashableTest extends FunSuite with GeneratorDrivenPropertyChecks { ) def testConsistency[T](algo: Hashable[Array[Byte], T]): Unit = { - forAll { input: Array[Byte] => - assert(algo(input) == algo(input)) - } + forAll { (input: Array[Byte]) => assert(algo(input) == algo(input)) } } algorithms foreach { algo => diff --git a/util-hashing/src/test/scala/com/twitter/hashing/KetamaDistributorTest.scala b/util-hashing/src/test/scala/com/twitter/hashing/KetamaDistributorTest.scala deleted file mode 100644 index 9c7f5e6620..0000000000 --- a/util-hashing/src/test/scala/com/twitter/hashing/KetamaDistributorTest.scala +++ /dev/null @@ -1,105 +0,0 @@ -package com.twitter.hashing - -import org.junit.runner.RunWith -import org.scalacheck.Gen -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks -import scala.collection.mutable -import java.io.{BufferedReader, InputStreamReader} -import java.nio.{ByteBuffer, ByteOrder} -import java.security.MessageDigest - -@RunWith(classOf[JUnitRunner]) -class KetamaDistributorTest extends WordSpec with GeneratorDrivenPropertyChecks { - "KetamaDistributor" should { - val nodes = Seq( - KetamaNode("10.0.1.1", 600, 1), - KetamaNode("10.0.1.2", 300, 2), - KetamaNode("10.0.1.3", 200, 3), - KetamaNode("10.0.1.4", 350, 4), - KetamaNode("10.0.1.5", 1000, 5), - KetamaNode("10.0.1.6", 800, 6), - KetamaNode("10.0.1.7", 950, 7), - KetamaNode("10.0.1.8", 100, 8) - ) - - // 160 is the hard coded value for libmemcached, which was this input data is from - val ketamaDistributor = new KetamaDistributor(nodes, 160) - val ketamaDistributorInoldLibMemcachedVersionComplianceMode = - new KetamaDistributor(nodes, 160, true) - "pick the correct node with ketama hash function" in { - // Test from Smile's KetamaNodeLocatorSpec.scala - - // Load known good results (key, hash(?), continuum ceiling(?), IP) - val stream = getClass.getClassLoader.getResourceAsStream("ketama_results") - val reader = new BufferedReader(new InputStreamReader(stream)) - val expected = new mutable.ListBuffer[Array[String]] - var line: String = null - do { - line = reader.readLine - if (line != null) { - val segments = line.split(" ") - assert(segments.length == 4) - expected += segments - } - } while (line != null) - assert(expected.size == 99) - - // Test that ketamaClient.clientOf(key) == expected IP - val handleToIp = nodes.map { n => - n.handle -> n.identifier - }.toMap - for (testcase <- expected) { - val hash = KeyHasher.KETAMA.hashKey(testcase(0).getBytes) - - val handle = ketamaDistributor.nodeForHash(hash) - val handle2 = ketamaDistributorInoldLibMemcachedVersionComplianceMode.nodeForHash(hash) - val resultIp = handleToIp(handle) - val resultIp2 = handleToIp(handle2) - assert(testcase(3) == resultIp) - assert(testcase(3) == resultIp2) - } - } - - "pick the correct node with 64-bit hash values" in { - val knownGoodValues = Map( - -166124121512512L -> 5, - 8796093022208L -> 3, - 4312515125124L -> 2, - -8192481414141L -> 1, - -9515121512312L -> 5 - ) - - knownGoodValues foreach { - case (key, node) => - val handle = ketamaDistributor.nodeForHash(key) - assert(handle == node) - } - } - - "hashInt" in { - val ketama = new KetamaDistributor[Unit](Seq.empty, 0, false) - forAll(Gen.chooseNum(0, Int.MaxValue)) { i => - val md = MessageDigest.getInstance("MD5") - ketama.hashInt(i, md) - val array = md.digest() - - assert(array.deep == md.digest(i.toString.getBytes("UTF-8")).deep) - } - } - - "byteArrayToLE" in { - val ketama = new KetamaDistributor[Unit](Seq.empty, 0, false) - forAll { s: String => - val ba = MessageDigest.getInstance("MD5").digest(s.getBytes("UTF-8")) - val bb = ByteBuffer.wrap(ba).order(ByteOrder.LITTLE_ENDIAN) - - assert(ketama.byteArrayToLE(ba, 0) == bb.getInt(0)) - assert(ketama.byteArrayToLE(ba, 4) == bb.getInt(4)) - assert(ketama.byteArrayToLE(ba, 8) == bb.getInt(8)) - assert(ketama.byteArrayToLE(ba, 12) == bb.getInt(12)) - } - } - } -} diff --git a/util-hashing/src/test/scala/com/twitter/hashing/KeyHasherTest.scala b/util-hashing/src/test/scala/com/twitter/hashing/KeyHasherTest.scala index 27cd1f67ce..57569ae3e4 100644 --- a/util-hashing/src/test/scala/com/twitter/hashing/KeyHasherTest.scala +++ b/util-hashing/src/test/scala/com/twitter/hashing/KeyHasherTest.scala @@ -2,18 +2,15 @@ package com.twitter.hashing import com.twitter.io.TempFile import java.util.Base64 -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner import scala.collection.mutable.ListBuffer import java.nio.charset.StandardCharsets.UTF_8 +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class KeyHasherTest extends WordSpec { +class KeyHasherTest extends AnyFunSuite { def readResource(name: String) = { var lines = new ListBuffer[String]() val src = scala.io.Source.fromFile(TempFile.fromResourcePath(getClass, "/" + name)) - src.getLines + src.getLines() } val base64 = Base64.getDecoder @@ -31,30 +28,27 @@ class KeyHasherTest extends WordSpec { } } - "KeyHasher" should { - "correctly hash fnv1_32" in { - testHasher("fnv1_32", KeyHasher.FNV1_32) - } - - "correctly hash fnv1_64" in { - testHasher("fnv1_64", KeyHasher.FNV1_64) - } + test("KeyHasher should correctly hash fnv1_32") { + testHasher("fnv1_32", KeyHasher.FNV1_32) + } - "correctly hash fnv1a_32" in { - testHasher("fnv1a_32", KeyHasher.FNV1A_32) - } + test("KeyHasher should correctly hash fnv1_64") { + testHasher("fnv1_64", KeyHasher.FNV1_64) + } - "correctly hash fnv1a_64" in { - testHasher("fnv1a_64", KeyHasher.FNV1A_64) - } + test("KeyHasher should correctly hash fnv1a_32") { + testHasher("fnv1a_32", KeyHasher.FNV1A_32) + } - "correctly hash jenkins" in { - testHasher("jenkins", KeyHasher.JENKINS) - } + test("KeyHasher should correctly hash fnv1a_64") { + testHasher("fnv1a_64", KeyHasher.FNV1A_64) + } - "correctly hash crc32 itu" in { - testHasher("crc32", KeyHasher.CRC32_ITU) - } + test("KeyHasher should correctly hash jenkins") { + testHasher("jenkins", KeyHasher.JENKINS) + } + test("KeyHasher should correctly hash crc32 itu") { + testHasher("crc32", KeyHasher.CRC32_ITU) } } diff --git a/util-intellij/OWNERS b/util-intellij/OWNERS deleted file mode 100644 index 80c8af4f73..0000000000 --- a/util-intellij/OWNERS +++ /dev/null @@ -1,6 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta -slance diff --git a/util-intellij/README.md b/util-intellij/README.md deleted file mode 100644 index 86831a9945..0000000000 --- a/util-intellij/README.md +++ /dev/null @@ -1,103 +0,0 @@ -# util-intellij - -`util-intellij` is a project containing IntelliJ plugins and other utilities. - -## Contents - -### Twitter Futures Capture Points - -Stack traces produced from asynchronous execution have undesirable properties, -e.g., they are disorganized and often contain frames from internal libraries -that programmers need not know about. This result is inevitable since not only -do different parts of the same program get executed on different threads, and -possibly different cores, but execution also jumps between different frames. -This means that each thread used to execute a program may produce a unique stack -trace. Thus, an Asynchronous Stack Trace (AST) gives a fragmented window view -into the execution of a program and thus makes it difficult to understand the -flow of code because the causality between frames is lost. - -In response to this phenomenon, IntelliJ 2017.1 introduced a feature called -Capture Points which is built on top of the IntelliJ Debugger. A capture point -is a place in a computer program where the debugger collects and saves stack -frames to be used later when we reach a specific point in the code and we want -to see how we got there. IntelliJ IDEA does this by substituting part of the -call stack with the captured frame. - -We have written capture points for Twitter Futures in an XML file. Users can -debug their asynchronous code more easily with these capture points. Note that -as of October 2017, the capture points work with only Scala 2.11.11. - -#### Setup - -The capture points are defined in -`util/util-intellij/src/main/resources/TwitterFuturesCapturePoints.xml`. -To import them into IntelliJ, - -1. Open IntelliJ. In the menu bar, click IntelliJ IDEA > Preferences. -2. Navigate to Build, Execution, Deployment > Debugger > Async Stacktraces. -3. Click the Import icon on the bottom bar. Find the XML file and click OK. - -#### Use - -TL;DR: set a breakpoint where you wish to see the stack trace, debug your code, -and look at the "Frames" tab in the Debugger. Any asynchronous calls in the -stack trace will appear in logical order. If you wish to clean up the stack -trace, click the funnel icon in the top right, "Hide Frames from Libraries". - -We will illustrate how to use the capture points to assist with debugging with a -small example. The example is located in -`util/util-intellij/src/test/scala/com/twitter/util/capturepoints/Demo.scala`. - -A brief explanation of this test: the test passes a `Promise[Int]` through three -methods and then sets the `Promise`’s value in a `futurePool`. The calls are -asynchronous, but the logical flow of the test is as follows: - -1. `test` block calls `someBusinessLogic` -2. `someBusinessLogic` calls `moreBusinessLogic` -3. `moreBusinessLogic` calls `lordBusinessLogic` -4. `lordBusinessLogic` waits for the `Promise`’s value to be set -5. The `Promise`’s value is set in the test block (this could happen at any - time; it is not necessarily step number 5) -6. `lordBusinessLogic` returns -7. `test`’s `result` variable is `4`, and the test passes - -Suppose we wish to inspect the stack trace from inside `lordBusinessLogic`. If -we set a breakpoint at line 47 and then debug the test with pants in IntelliJ -with the capture points disabled, we see the following stack trace: - -``` -apply$mcVI$sp:47, Demo$$anonfun$lordBusinessLogic$1 -apply:1820, Future$$anonfun$onSuccess$1 -apply:205, Promise$Monitored -run:532, Promise$$anon$7 -run:198, LocalScheduler$Activation -submit:157, LocalScheduler$Activation -submit:274, LocalScheduler -submit:109, Scheduler$ -runq:522, Promise -updatelfEmpty:880, Promise -update:859, Promise -setValue:835, Promise -apply$mcV$sp:20, Demo$$anonfun$3$$anonfun$apply$1 -apply:15, Try$ -run:140, ExecutorServiceFuturePool$$anon$4 -``` - -The library calls in the stack trace do not help us understand our code. If we -enable the capture points and again debug the test, we see the following stack -trace: - -``` -apply$mcVI$sp:47, Demo$$anonfun$lordBusinessLogic$1 -apply:1820, Future$$anonfun$onSuccess$1 -onSuccess:1819, Future -lordBusinessLogic:42, Demo -moreBusinessLogic:38, Demo -someBusinessLogic:31, Demo -apply:17, Demo$$anonfun$3 -``` - -This is much cleaner and it includes only calls to functions we wrote. It lets -us clearly see that the test block called `someBusinessLogic`, that -`someBusinessLogic` called `moreBusinessLogic`, and that `moreBusinessLogic` -called `lordBusinessLogic`. diff --git a/util-intellij/src/main/resources/BUILD b/util-intellij/src/main/resources/BUILD deleted file mode 100644 index cc44b5fc88..0000000000 --- a/util-intellij/src/main/resources/BUILD +++ /dev/null @@ -1,3 +0,0 @@ -resources( - sources = rglobs("*"), -) diff --git a/util-intellij/src/main/resources/util-intellij/TwitterFuturesCapturePoints.xml b/util-intellij/src/main/resources/util-intellij/TwitterFuturesCapturePoints.xml deleted file mode 100644 index 155a33a89a..0000000000 --- a/util-intellij/src/main/resources/util-intellij/TwitterFuturesCapturePoints.xml +++ /dev/null @@ -1,138 +0,0 @@ - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - diff --git a/util-intellij/src/test/scala/com/twitter/util/capturepoints/BUILD b/util-intellij/src/test/scala/com/twitter/util/capturepoints/BUILD deleted file mode 100644 index 2bcee27eb7..0000000000 --- a/util-intellij/src/test/scala/com/twitter/util/capturepoints/BUILD +++ /dev/null @@ -1,10 +0,0 @@ -junit_tests( - sources = globs("*.scala"), - fatal_warnings = True, - strict_deps = True, - dependencies = [ - "3rdparty/jvm/org/scalatest", - "util/util-core", - "util/util-intellij/src/main/resources", - ], -) diff --git a/util-intellij/src/test/scala/com/twitter/util/capturepoints/Demo.scala b/util-intellij/src/test/scala/com/twitter/util/capturepoints/Demo.scala deleted file mode 100644 index 1b006ec4eb..0000000000 --- a/util-intellij/src/test/scala/com/twitter/util/capturepoints/Demo.scala +++ /dev/null @@ -1,60 +0,0 @@ -package com.twitter.util.capturepoints - -import com.twitter.util.{Await, Future, FuturePool, Promise} -import org.scalatest.FunSuite - -/** - * A trivial unit test used for demonstrating the value of async stacktraces. - * - * Place a breakpoint in the `onSuccess` method inside `lordBusinessLogic` - * to see the difference in the actual stacktrace versus the IDE's synthetic - * one. - */ -class Demo extends FunSuite { - - test("async stack traces are rad") { - val promise = new Promise[Int]() - val future = someBusinessLogic(promise) - - FuturePool.unboundedPool { - promise.setValue(1) - } - - val result: Int = Await.result(future) - assert(4 == result) - } - - def someBusinessLogic(future: Future[Int]): Future[Int] = { - val result: Future[Int] = future.map { i => - i + 1 - } - moreBusinessLogic(result) - } - - def moreBusinessLogic(future: Future[Int]): Future[Int] = { - val result: Future[Int] = future.map { i => - i * 2 - } - lordBusinessLogic(result) - } - - def lordBusinessLogic(future: Future[Int]): Future[Int] = { - future.onSuccess { i => - // NOTE: with Intellij's async stacktraces configured a breakpoint - // here *will include* `someBusinessLogic` and `moreBusinessLogic` - // in the debugger's stacktrace, despite them not being in the - // actual stacktrace. - val stackTrace = Thread.currentThread.getStackTrace.mkString("\n") - println(stackTrace) - - assert(!stackTrace.contains("someBusinessLogic")) - assert(!stackTrace.contains("moreBusinessLogic")) - assert(stackTrace.contains("lordBusinessLogic")) - - println("\n (⊙_◎) (⧉ ⦣ ⧉) (`_´)ゞ 「(°ヘ°)\n") - println(s" How did we get here? Where are `someBusinessLogic` and `moreBusinessLogic`?") - println("") - } - } - -} diff --git a/util-intellij/src/test/scala/com/twitter/util/capturepoints/TwitterFuturesCapturePointsTest.scala b/util-intellij/src/test/scala/com/twitter/util/capturepoints/TwitterFuturesCapturePointsTest.scala deleted file mode 100644 index 7145f70f59..0000000000 --- a/util-intellij/src/test/scala/com/twitter/util/capturepoints/TwitterFuturesCapturePointsTest.scala +++ /dev/null @@ -1,54 +0,0 @@ -package com.twitter.util.capturepoints - -import java.lang.reflect.Method -import org.scalatest.FunSuite -import scala.xml.XML - -/** - * This test suite is intended to test capture points defined in [TwitterFuturesCapturePoints.xml]. - * The test name corresponds to the capture point being tested and is named in the format - * [Class.method]. Because of the limitations of reflection, we only test for [Capture class - * and method] and [Insert class and method]. This means that the [key-expression] is not - * tested for. As such, these tests will only fail if the class name or method for which a - * capture point depends on has either changed or doesn't exist anymore. - * - * In the case that the [Insert class-name] or [Capture class-name] has been changed or - * doesn't exist, this file will fail to compile. - */ -class TwitterFuturesCapturePointsTest extends FunSuite { - - /** - * extract [methodName] from `Class[_]` - */ - private[this] def findMethod(m: Class[_], methodName: String = "apply"): Array[Method] = { - m.getDeclaredMethods.filter(_.getName == methodName) - } - - /** - * test that [methodName] is a member of `ClassOf[com.twitter.util.className]` - */ - private[this] def testMethod(methodType: String, className: String, methodName: String = "apply"): Unit = { - val m = Class.forName(className) - val methodArray = findMethod(m, methodName) - assert(methodArray.nonEmpty, s": couldn't find $methodType method ${m.getTypeName}.$methodName.") - } - - val capturePointsXML = XML.load(getClass.getResourceAsStream("/util-intellij/TwitterFuturesCapturePoints.xml")) - - /** - * perform "capture method" and "insert method" tests for each capture point - */ - (capturePointsXML \ "capture-point").foreach { capturePoint => - val attrMap = capturePoint.attributes.asAttrMap - val insertMethodName = attrMap("insert-method-name") - val insertClassName = attrMap("insert-class-name") - val captureMethodName = attrMap("method-name") - val captureClassName = attrMap("class-name") - val testName = s"$captureClassName.$captureMethodName" - - test(testName) { - testMethod("capture", captureClassName, captureMethodName) - testMethod("insert", insertClassName, insertMethodName) - } - } -} diff --git a/util-jackson-annotations/PROJECT b/util-jackson-annotations/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-jackson-annotations/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/BUILD b/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/BUILD new file mode 100644 index 0000000000..0b686a99bb --- /dev/null +++ b/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/BUILD @@ -0,0 +1,12 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-annotations", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/InjectableValue.java b/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/InjectableValue.java new file mode 100644 index 0000000000..5c8a353697 --- /dev/null +++ b/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation/InjectableValue.java @@ -0,0 +1,19 @@ +package com.twitter.util.jackson.annotation; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * Marker annotation for {@link java.lang.annotation.Annotation} interfaces which define + * a field that should be injected into a case class via Jackson + * `com.fasterxml.jackson.databind.InjectableValues`. + * + * @see com.fasterxml.jackson.databind.InjectableValues + * @see Finatra User's Guide - JSON Injectable Values + * @see Finatra User's Guide - HTTP Request Field Annotations + */ +@Target(ElementType.ANNOTATION_TYPE) +@Retention(RetentionPolicy.RUNTIME) +public @interface InjectableValue {} diff --git a/util-jackson/PROJECT b/util-jackson/PROJECT new file mode 100644 index 0000000000..c37ef6c2b7 --- /dev/null +++ b/util-jackson/PROJECT @@ -0,0 +1,8 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan + - tapans +watchers: + - csl-team@twitter.com diff --git a/util-jackson/src/main/scala/com/fasterxml/jackson/databind/BUILD b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/BUILD new file mode 100644 index 0000000000..d82b398136 --- /dev/null +++ b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/BUILD @@ -0,0 +1,20 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-objectmapper-copy", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + ], + exports = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + ], +) diff --git a/util-jackson/src/main/scala/com/fasterxml/jackson/databind/DeserializationContextAccessor.scala b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/DeserializationContextAccessor.scala new file mode 100644 index 0000000000..dc6c032557 --- /dev/null +++ b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/DeserializationContextAccessor.scala @@ -0,0 +1,7 @@ +package com.fasterxml.jackson.databind + +/** INTERNAL API ONLY */ +class DeserializationContextAccessor(context: DeserializationContext) { + // Workaround to get access to the DeserializationContext#injectableValues member + def getInjectableValues: Option[InjectableValues] = Option(context._injectableValues) +} diff --git a/util-jackson/src/main/scala/com/fasterxml/jackson/databind/ObjectMapperCopier.scala b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/ObjectMapperCopier.scala new file mode 100644 index 0000000000..369175d89e --- /dev/null +++ b/util-jackson/src/main/scala/com/fasterxml/jackson/databind/ObjectMapperCopier.scala @@ -0,0 +1,11 @@ +package com.fasterxml.jackson.databind + +import com.fasterxml.jackson.module.scala.ScalaObjectMapper + +/** INTERNAL API ONLY */ +object ObjectMapperCopier { + // Workaround since calling objectMapper.copy on "ObjectMapper with ScalaObjectMapper" fails the _checkInvalidCopy check + def copy(objectMapper: ObjectMapper): ObjectMapper with ScalaObjectMapper = { + new ObjectMapper(objectMapper) with ScalaObjectMapper + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/BUILD b/util-jackson/src/main/scala/com/twitter/util/jackson/BUILD new file mode 100644 index 0000000000..fdcebdf661 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/BUILD @@ -0,0 +1,38 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-core", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/dataformat:jackson-dataformat-yaml", + "3rdparty/jvm/com/fasterxml/jackson/datatype:jackson-datatype-jsr310", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-jackson/src/main/scala/com/fasterxml/jackson/databind", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/serde", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-validator/src/main/scala/com/twitter/util/validation", + ], + exports = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/dataformat:jackson-dataformat-yaml", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + ], +) diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/JSON.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/JSON.scala new file mode 100644 index 0000000000..2008cbf22b --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/JSON.scala @@ -0,0 +1,69 @@ +package com.twitter.util.jackson + +import com.twitter.io.{Buf, ClasspathResource} +import com.twitter.util.Try +import java.io.{File, InputStream} + +/** + * Uses an instance of a [[ScalaObjectMapper]] configured to not perform any type of validation. + * Inspired by [[https://www.scala-lang.org/api/2.12.6/scala-parser-combinators/scala/util/parsing/json/JSON$.html]]. + * + * @note This is only intended for use from Scala (not Java). + * + * @see [[ScalaObjectMapper]] + */ +object JSON { + private final val Mapper: ScalaObjectMapper = + ScalaObjectMapper.builder.withNoValidation.objectMapper + + /** Simple utility to parse a JSON string into an Option[T] type. */ + def parse[T: Manifest](input: String): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** Simple utility to parse a JSON [[Buf]] into an Option[T] type. */ + def parse[T: Manifest](input: Buf): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** + * Simple utility to parse a JSON [[InputStream]] into an Option[T] type. + * @note the caller is responsible for managing the lifecycle of the given [[InputStream]]. + */ + def parse[T: Manifest](input: InputStream): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** + * Simple utility to load a JSON file and parse contents into an Option[T] type. + * + * @note the caller is responsible for managing the lifecycle of the given [[File]]. + */ + def parse[T: Manifest](f: File): Option[T] = + Try(Mapper.underlying.readValue[T](f)).toOption + + /** Simple utility to write a value as a JSON encoded String. */ + def write(any: Any): String = + Mapper.writeValueAsString(any) + + /** Simple utility to pretty print a JSON encoded String from the given instance. */ + def prettyPrint(any: Any): String = + Mapper.writePrettyString(any) + + object Resource { + + /** + * Simple utility to load a JSON resource from the classpath and parse contents into an Option[T] type. + * + * @note `name` resolution to locate the resource is governed by [[java.lang.Class#getResourceAsStream]] + */ + def parse[T: Manifest](name: String): Option[T] = { + ClasspathResource.load(name) match { + case Some(inputStream) => + try { + JSON.parse[T](inputStream) + } finally { + inputStream.close() + } + case _ => None + } + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/JsonDiff.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/JsonDiff.scala new file mode 100644 index 0000000000..cf2cf9edd5 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/JsonDiff.scala @@ -0,0 +1,229 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.annotation.JsonInclude.Include +import com.fasterxml.jackson.databind.node.TextNode +import com.fasterxml.jackson.databind.{JsonNode, ObjectMapperCopier, SerializationFeature} +import com.twitter.util.logging.Logging +import java.io.PrintStream +import scala.util.control.NonFatal + +object JsonDiff extends Logging { + + private[this] lazy val mapper: ScalaObjectMapper = + ScalaObjectMapper.builder.objectMapper + + private[this] lazy val sortingObjectMapper: JacksonScalaObjectMapperType = { + val newMapper = ObjectMapperCopier.copy(mapper.underlying) + newMapper.configure(SerializationFeature.ORDER_MAP_ENTRIES_BY_KEYS, true) + newMapper.setDefaultPropertyInclusion(Include.ALWAYS) + newMapper + } + + /** + * Creates a string representation of the given [[JsonNode]] with entries + * sorted alphabetically by key. + * + * @param jsonNode - input [[JsonNode]] + * @return string representation of the [[JsonNode]]. + */ + def toSortedString(jsonNode: JsonNode): String = { + if (jsonNode.isTextual) { + jsonNode.textValue() + } else { + val node = sortingObjectMapper.treeToValue(jsonNode, classOf[Object]) + sortingObjectMapper.writeValueAsString(node) + } + } + + /** A [[JsonDiff]] result */ + case class Result( + expected: JsonNode, + expectedPrettyString: String, + actual: JsonNode, + actualPrettyString: String) { + + private[this] val expectedJsonSorted = JsonDiff.toSortedString(expected) + private[this] val actualJsonSorted = JsonDiff.toSortedString(actual) + + override val toString: String = { + val expectedHeader = "Expected: " + val diffStartIdx = + actualJsonSorted.zip(expectedJsonSorted).indexWhere { case (x, y) => x != y } + + val message = new StringBuilder + message.append(" " * (expectedHeader.length + diffStartIdx) + "*\n") + message.append(expectedHeader + expectedJsonSorted + "\n") + message.append(s"Actual: $actualJsonSorted") + + message.toString() + } + } + + /** + * Computes the diff for two snippets of json both of expected type `T`. + * If a difference is detected a [[Result]] is returned otherwise a None. + * + * @param expected the expected json + * @param actual the actual or received json + * @return if a difference is detected a [[Result]] + * is returned otherwise a None. + */ + def diff[T](expected: T, actual: T): Option[Result] = + diff[T](expected, actual, None) + + /** + * Computes the diff for two snippets of json both of expected type `T`. + * If a difference is detected a [[Result]] is returned otherwise a None. + * + * @param expected the expected json + * @param actual the actual or received json + * @param normalizeFn a function to apply to the actual json in order to "normalize" values. + * + * @return if a difference is detected a [[Result]] + * is returned otherwise a None. + * + * ==Usage== + * {{{ + * + * private def normalize(jsonNode: JsonNode): JsonNode = jsonNode match { + * case on: ObjectNode => on.put("time", "1970-01-01T00:00:00Z") + * case _ => jsonNode + * } + * + * val expected = """{"foo": "bar", "time": ""1970-01-01T00:00:00Z"}""" + * val actual = ??? ({"foo": "bar", "time": ""2021-05-14T00:00:00Z"}) + * val result: Option[JsonDiff.Result] = JsonDiff.diff(expected, actual, normalize) + * }}} + */ + def diff[T]( + expected: T, + actual: T, + normalizeFn: JsonNode => JsonNode + ): Option[Result] = + diff[T](expected, actual, Option(normalizeFn)) + + /** + * Asserts that the actual equals the expected. Will throw an [[AssertionError]] with details + * printed to [[System.out]]. + * + * @param expected the expected json + * @param actual the actual or received json + * + * @throws AssertionError - when the expected does not match the actual. + */ + @throws[AssertionError] + def assertDiff[T](expected: T, actual: T): Unit = + assert[T](expected, actual, None, System.out) + + /** + * Asserts that the actual equals the expected. Will throw an [[AssertionError]] with details + * printed to [[System.out]] using the given normalizeFn to normalize the actual contents. + * + * @param expected the expected json + * @param actual the actual or received json + * @param normalizeFn a function to apply to the actual json in order to "normalize" values. + * + * @throws AssertionError - when the expected does not match the actual. + */ + @throws[AssertionError] + def assertDiff[T]( + expected: T, + actual: T, + normalizeFn: JsonNode => JsonNode + ): Unit = assert[T](expected, actual, Option(normalizeFn), System.out) + + /** + * Asserts that the actual equals the expected. Will throw an [[AssertionError]] with details + * printed to the given [[PrintStream]]. + * + * @param expected the expected json + * @param actual the actual or received json + * @param p the [[PrintStream]] for reporting details + * + * @throws AssertionError - when the expected does not match the actual. + */ + @throws[AssertionError] + def assertDiff[T](expected: T, actual: T, p: PrintStream): Unit = + assert[T](expected, actual, None, p) + + /** + * Asserts that the actual equals the expected. Will throw an [[AssertionError]] with details + * printed to the given [[PrintStream]] using the given normalizeFn to normalize the actual contents. + * + * @param expected the expected json + * @param actual the actual or received json + * @param p the [[PrintStream]] for reporting details + * @param normalizeFn a function to apply to the actual json in order to "normalize" values. + * + * @throws AssertionError - when the expected does not match the actual. + */ + @throws[AssertionError] + def assertDiff[T]( + expected: T, + actual: T, + normalizeFn: JsonNode => JsonNode, + p: PrintStream + ): Unit = assert[T](expected, actual, Option(normalizeFn), p) + + /* Private */ + + private[this] def diff[T]( + expected: T, + actual: T, + normalizer: Option[JsonNode => JsonNode] + ): Option[Result] = { + val actualJson = jsonString(actual) + val expectedJson = jsonString(expected) + + val actualJsonNode: JsonNode = { + val jsonNode = tryJsonNodeParse(actualJson) + normalizer match { + case Some(normalizerFn) => normalizerFn(jsonNode) + case _ => jsonNode + } + } + + val expectedJsonNode: JsonNode = tryJsonNodeParse(expectedJson) + if (actualJsonNode != expectedJsonNode) { + val result = Result( + expected = expectedJsonNode, + expectedPrettyString = mapper.writePrettyString(expectedJsonNode), + actual = actualJsonNode, + actualPrettyString = mapper.writePrettyString(actualJsonNode) + ) + Some(result) + } else None + } + + private[this] def assert[T]( + expected: T, + actual: T, + normalizer: Option[JsonNode => JsonNode], + p: PrintStream + ): Unit = { + diff(expected, actual, normalizer) match { + case Some(result) => + p.println("JSON DIFF FAILED!") + p.println(result) + throw new AssertionError(s"${JsonDiff.getClass.getName} failure\n$result") + case _ => // do nothing + } + } + + private[this] def tryJsonNodeParse(expectedJsonStr: String): JsonNode = { + try { + mapper.parse[JsonNode](expectedJsonStr) + } catch { + case NonFatal(e) => + warn(e.getMessage) + new TextNode(expectedJsonStr) + } + } + + private[this] def jsonString(receivedJson: Any): String = { + receivedJson match { + case str: String => str + case _ => mapper.writeValueAsString(receivedJson) + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/ScalaObjectMapper.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/ScalaObjectMapper.scala new file mode 100644 index 0000000000..496122b164 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/ScalaObjectMapper.scala @@ -0,0 +1,647 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.annotation.JsonInclude +import com.fasterxml.jackson.annotation.JsonInclude.Include +import com.fasterxml.jackson.core.json.JsonWriteFeature +import com.fasterxml.jackson.core.util.DefaultIndenter +import com.fasterxml.jackson.core.util.DefaultPrettyPrinter +import com.fasterxml.jackson.core.JsonFactory +import com.fasterxml.jackson.core.JsonFactoryBuilder +import com.fasterxml.jackson.core.JsonParser +import com.fasterxml.jackson.core.TSFBuilder +import com.fasterxml.jackson.databind.util.ByteBufferBackedInputStream +import com.fasterxml.jackson.databind.{ObjectMapper => JacksonObjectMapper, _} +import com.fasterxml.jackson.dataformat.yaml.YAMLFactory +import com.fasterxml.jackson.dataformat.yaml.YAMLFactoryBuilder +import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule +import com.fasterxml.jackson.module.scala.DefaultScalaModule +import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} +import com.twitter.io.Buf +import com.twitter.util.jackson.caseclass.CaseClassJacksonModule +import com.twitter.util.jackson.serde.DefaultSerdeModule +import com.twitter.util.jackson.serde.LongKeyDeserializers +import com.twitter.util.validation.ScalaValidator +import java.io.ByteArrayOutputStream +import java.io.InputStream +import java.io.OutputStream +import java.nio.ByteBuffer + +object ScalaObjectMapper { + + /** The default [[ScalaValidator]] for a [[ScalaObjectMapper]] */ + private[twitter] val DefaultValidator: ScalaValidator = ScalaValidator() + + /** The default [[JsonWriteFeature.WRITE_NUMBERS_AS_STRINGS]] setting */ + private[twitter] val DefaultNumbersAsStrings: Boolean = false + + /** Framework modules need to be added 'last' so they can override existing ser/des */ + private[twitter] val DefaultJacksonModules: Seq[Module] = + Seq(DefaultScalaModule, new JavaTimeModule, LongKeyDeserializers, DefaultSerdeModule) + + /** The default [[PropertyNamingStrategy]] for a [[ScalaObjectMapper]] */ + private[twitter] val DefaultPropertyNamingStrategy: PropertyNamingStrategy = + PropertyNamingStrategies.SNAKE_CASE + + /** The default [[JsonInclude.Include]] for serialization for a [[ScalaObjectMapper]] */ + private[twitter] val DefaultSerializationInclude: JsonInclude.Include = + JsonInclude.Include.NON_ABSENT + + /** The default configuration for serialization as a `Map[SerializationFeature, Boolean]` */ + private[twitter] val DefaultSerializationConfig: Map[SerializationFeature, Boolean] = + Map( + SerializationFeature.WRITE_DATES_AS_TIMESTAMPS -> false, + SerializationFeature.WRITE_ENUMS_USING_TO_STRING -> true) + + /** The default configuration for deserialization as a `Map[DeserializationFeature, Boolean]` */ + private[twitter] val DefaultDeserializationConfig: Map[DeserializationFeature, Boolean] = + Map( + DeserializationFeature.FAIL_ON_NULL_FOR_PRIMITIVES -> true, + DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES -> false, + DeserializationFeature.READ_ENUMS_USING_TO_STRING -> true, + DeserializationFeature.FAIL_ON_READING_DUP_TREE_KEY -> true, + DeserializationFeature.USE_JAVA_ARRAY_FOR_JSON_ARRAY -> true /* see jackson-module-scala/issues/148 */ + ) + + /** The default for configuring additional modules on the underlying [[JacksonScalaObjectMapperType]] */ + private[twitter] val DefaultAdditionalJacksonModules: Seq[Module] = Seq.empty[Module] + + /** The default setting to enable case class validation during case class deserialization */ + private[twitter] val DefaultValidation: Boolean = true + + /** + * Returns a new [[ScalaObjectMapper]] configured with [[Builder]] defaults. + * + * @return a new [[ScalaObjectMapper]] instance. + * + * @see [[com.fasterxml.jackson.databind.InjectableValues]] + */ + def apply(): ScalaObjectMapper = + ScalaObjectMapper.builder.objectMapper + + /** + * Creates a new [[ScalaObjectMapper]] from an underlying [[JacksonScalaObjectMapperType]]. + * + * @note this mutates the underlying mapper to configure it with the [[ScalaObjectMapper]] defaults. + * + * @param underlying the [[JacksonScalaObjectMapperType]] to wrap. + * + * @return a new [[ScalaObjectMapper]] + */ + def apply(underlying: JacksonScalaObjectMapperType): ScalaObjectMapper = { + new ScalaObjectMapper( + ScalaObjectMapper.builder + .configureJacksonScalaObjectMapper(underlying) + ) + } + + /** + * Utility to create a new [[ScalaObjectMapper]] which simply wraps the given + * [[JacksonScalaObjectMapperType]]. + * + * @note the `underlying` mapper is not mutated to produce the new [[ScalaObjectMapper]] + */ + def objectMapper(underlying: JacksonScalaObjectMapperType): ScalaObjectMapper = { + val objectMapperCopy = ObjectMapperCopier.copy(underlying) + new ScalaObjectMapper(objectMapperCopy) + } + + /** + * Utility to create a new [[ScalaObjectMapper]] explicitly configured to serialize and deserialize + * YAML using the given [[JacksonScalaObjectMapperType]]. The resultant mapper [[PropertyNamingStrategy]] + * will be that configured on the `underlying` mapper. + * + * @note the `underlying` mapper is copied (not mutated) to produce the new [[ScalaObjectMapper]] + * to negotiate YAML serialization and deserialization. + */ + def yamlObjectMapper(underlying: JacksonScalaObjectMapperType): ScalaObjectMapper = + underlying.getFactory match { + case _: YAMLFactory => // correct + val objectMapperCopy = ObjectMapperCopier.copy(underlying) + new ScalaObjectMapper(objectMapperCopy) + case _ => // incorrect + throw new IllegalArgumentException( + s"The underlying mapper is not properly configured with a YAMLFactory") + } + + /** + * Utility to create a new [[ScalaObjectMapper]] explicitly configured with + * [[PropertyNamingStrategies.LOWER_CAMEL_CASE]] as a `PropertyNamingStrategy` wrapping the + * given [[JacksonScalaObjectMapperType]]. + * + * @note the `underlying` mapper is copied (not mutated) to produce the new [[ScalaObjectMapper]] + * with a [[PropertyNamingStrategies.LOWER_CAMEL_CASE]] PropertyNamingStrategy. + */ + def camelCaseObjectMapper(underlying: JacksonScalaObjectMapperType): ScalaObjectMapper = { + val objectMapperCopy = ObjectMapperCopier.copy(underlying) + objectMapperCopy.setPropertyNamingStrategy(PropertyNamingStrategies.LOWER_CAMEL_CASE) + new ScalaObjectMapper(objectMapperCopy) + } + + /** + * Utility to create a new [[ScalaObjectMapper]] explicitly configured with + * [[PropertyNamingStrategies.SNAKE_CASE]] as a `PropertyNamingStrategy` wrapping the + * given [[JacksonScalaObjectMapperType]]. + * + * @note the `underlying` mapper is copied (not mutated) to produce the new [[ScalaObjectMapper]] + * with a [[PropertyNamingStrategies.SNAKE_CASE]] PropertyNamingStrategy. + */ + def snakeCaseObjectMapper(underlying: JacksonScalaObjectMapperType): ScalaObjectMapper = { + val objectMapperCopy = ObjectMapperCopier.copy(underlying) + objectMapperCopy.setPropertyNamingStrategy(PropertyNamingStrategies.SNAKE_CASE) + new ScalaObjectMapper(objectMapperCopy) + } + + /** + * + * Build a new instance of a [[ScalaObjectMapper]]. + * + * For example, + * {{{ + * ScalaObjectMapper.builder + * .withPropertyNamingStrategy(new PropertyNamingStrategies.UpperCamelCaseStrategy) + * .withNumbersAsStrings(true) + * .withAdditionalJacksonModules(...) + * .objectMapper + * }}} + * + * or + * + * {{{ + * val builder = + * ScalaObjectMapper.builder + * .withPropertyNamingStrategy(new PropertyNamingStrategies.UpperCamelCaseStrategy) + * .withNumbersAsStrings(true) + * .withAdditionalJacksonModules(...) + * + * val mapper = builder.objectMapper + * val camelCaseMapper = builder.camelCaseObjectMapper + * }}} + * + */ + def builder: ScalaObjectMapper.Builder = Builder() + + /** + * A Builder for creating a new [[ScalaObjectMapper]]. E.g., to build a new instance of + * a [[ScalaObjectMapper]]. + * + * For example, + * {{{ + * ScalaObjectMapper.builder + * .withPropertyNamingStrategy(new PropertyNamingStrategies.UpperCamelCaseStrategy) + * .withNumbersAsStrings(true) + * .withAdditionalJacksonModules(...) + * .objectMapper + * }}} + * + * or + * + * {{{ + * val builder = + * ScalaObjectMapper.builder + * .withPropertyNamingStrategy(new PropertyNamingStrategies.UpperCamelCaseStrategy) + * .withNumbersAsStrings(true) + * .withAdditionalJacksonModules(...) + * + * val mapper = builder.objectMapper + * val camelCaseMapper = builder.camelCaseObjectMapper + * }}} + */ + case class Builder private[jackson] ( + propertyNamingStrategy: PropertyNamingStrategy = DefaultPropertyNamingStrategy, + numbersAsStrings: Boolean = DefaultNumbersAsStrings, + serializationInclude: Include = DefaultSerializationInclude, + serializationConfig: Map[SerializationFeature, Boolean] = DefaultSerializationConfig, + deserializationConfig: Map[DeserializationFeature, Boolean] = DefaultDeserializationConfig, + defaultJacksonModules: Seq[Module] = DefaultJacksonModules, + validator: Option[ScalaValidator] = Some(DefaultValidator), + additionalJacksonModules: Seq[Module] = DefaultAdditionalJacksonModules, + additionalMapperConfigurationFns: Seq[JacksonObjectMapper => Unit] = Seq.empty, + validation: Boolean = DefaultValidation) { + + /* Public */ + + /** Create a new [[ScalaObjectMapper]] from this [[Builder]]. */ + final def objectMapper: ScalaObjectMapper = + new ScalaObjectMapper(jacksonScalaObjectMapper) + + /** Create a new [[ScalaObjectMapper]] from this [[Builder]] using the given [[JsonFactory]]. */ + final def objectMapper[F <: JsonFactory](factory: F): ScalaObjectMapper = + new ScalaObjectMapper(jacksonScalaObjectMapper(factory)) + + /** + * Create a new [[ScalaObjectMapper]] explicitly configured to serialize and deserialize + * YAML from this [[Builder]]. + * + * @note the used [[PropertyNamingStrategy]] is defined by the current [[Builder]] configuration. + */ + final def yamlObjectMapper: ScalaObjectMapper = + new ScalaObjectMapper( + configureJacksonScalaObjectMapper(new YAMLFactoryBuilder(new YAMLFactory()))) + + /** + * Creates a new [[ScalaObjectMapper]] explicitly configured with + * [[PropertyNamingStrategies.LOWER_CAMEL_CASE]] as a `PropertyNamingStrategy`. + */ + final def camelCaseObjectMapper: ScalaObjectMapper = + ScalaObjectMapper.camelCaseObjectMapper(jacksonScalaObjectMapper) + + /** + * Creates a new [[ScalaObjectMapper]] explicitly configured with + * [[PropertyNamingStrategies.SNAKE_CASE]] as a `PropertyNamingStrategy`. + */ + final def snakeCaseObjectMapper: ScalaObjectMapper = + ScalaObjectMapper.snakeCaseObjectMapper(jacksonScalaObjectMapper) + + /* Builder Methods */ + + /** + * Configure a [[PropertyNamingStrategy]] for this [[Builder]]. + * @note the default is [[PropertyNamingStrategies.SNAKE_CASE]] + * @see [[ScalaObjectMapper.DefaultPropertyNamingStrategy]] + */ + final def withPropertyNamingStrategy(propertyNamingStrategy: PropertyNamingStrategy): Builder = + this.copy(propertyNamingStrategy = propertyNamingStrategy) + + /** + * Enable the [[JsonWriteFeature.WRITE_NUMBERS_AS_STRINGS]] for this [[Builder]]. + * @note the default is false. + */ + final def withNumbersAsStrings(numbersAsStrings: Boolean): Builder = + this.copy(numbersAsStrings = numbersAsStrings) + + /** + * Configure a [[JsonInclude.Include]] for serialization for this [[Builder]]. + * @note the default is [[JsonInclude.Include.NON_ABSENT]] + * @see [[ScalaObjectMapper.DefaultSerializationInclude]] + */ + final def withSerializationInclude(serializationInclude: Include): Builder = + this.copy(serializationInclude = serializationInclude) + + /** + * Set the serialization configuration for this [[Builder]] as a `Map` of `SerializationFeature` + * to `Boolean` (enabled). + * @note the default is described by [[ScalaObjectMapper.DefaultSerializationConfig]]. + * @see [[ScalaObjectMapper.DefaultSerializationConfig]] + */ + final def withSerializationConfig( + serializationConfig: Map[SerializationFeature, Boolean] + ): Builder = + this.copy(serializationConfig = serializationConfig) + + /** + * Set the deserialization configuration for this [[Builder]] as a `Map` of `DeserializationFeature` + * to `Boolean` (enabled). + * @note this overwrites the default deserialization configuration of this [[Builder]]. + * @note the default is described by [[ScalaObjectMapper.DefaultDeserializationConfig]]. + * @see [[ScalaObjectMapper.DefaultDeserializationConfig]] + */ + final def withDeserializationConfig( + deserializationConfig: Map[DeserializationFeature, Boolean] + ): Builder = + this.copy(deserializationConfig = deserializationConfig) + + /** + * Configure a [[ScalaValidator]] for this [[Builder]] + * @see [[ScalaObjectMapper.DefaultValidator]] + * + * @note If you pass `withNoValidation` to the builder all case class validations will be + * bypassed, regardless of the `withValidator` configuration. + */ + final def withValidator(validator: ScalaValidator): Builder = + this.copy(validator = Some(validator)) + + /** + * Configure the list of additional Jackson [[Module]]s for this [[Builder]]. + * @note this will overwrite (not append) the list additional Jackson [[Module]]s of this [[Builder]]. + */ + final def withAdditionalJacksonModules(additionalJacksonModules: Seq[Module]): Builder = + this.copy(additionalJacksonModules = additionalJacksonModules) + + /** + * Configure additional [[JacksonObjectMapper]] functionality for the underlying mapper of this [[Builder]]. + * @note this will overwrite any previously set function. + */ + final def withAdditionalMapperConfigurationFn(mapperFn: JacksonObjectMapper => Unit): Builder = + this + .copy(additionalMapperConfigurationFns = this.additionalMapperConfigurationFns :+ mapperFn) + + /** Method to allow changing of the default Jackson Modules for use from the `ScalaObjectMapperModule` */ + private[twitter] final def withDefaultJacksonModules( + defaultJacksonModules: Seq[Module] + ): Builder = + this.copy(defaultJacksonModules = defaultJacksonModules) + + /** + * Disable case class validation during case class deserialization + * + * @see [[ScalaObjectMapper.DefaultValidation]] + * @note If you pass `withNoValidation` to the builder all case class validations will be + * bypassed, regardless of the `withValidator` configuration. + */ + final def withNoValidation: Builder = + this.copy(validation = false) + + /* Private */ + + private[this] def defaultMapperConfiguration(mapper: JacksonObjectMapper): Unit = { + /* Serialization Config */ + mapper.setDefaultPropertyInclusion( + JsonInclude.Value.construct(serializationInclude, serializationInclude)) + mapper + .configOverride(classOf[Option[_]]) + .setIncludeAsProperty(JsonInclude.Value.construct(serializationInclude, Include.ALWAYS)) + for ((feature, state) <- serializationConfig) { + mapper.configure(feature, state) + } + + /* Deserialization Config */ + for ((feature, state) <- deserializationConfig) { + mapper.configure(feature, state) + } + } + + /** Order is important: default + case class module + any additional */ + private[this] def jacksonModules: Seq[Module] = { + this.defaultJacksonModules ++ + Seq(new CaseClassJacksonModule(if (this.validation) this.validator else None)) ++ + this.additionalJacksonModules + } + + private[this] final def jacksonScalaObjectMapper: JacksonScalaObjectMapperType = + configureJacksonScalaObjectMapper(new JsonFactoryBuilder) + + private[this] final def jacksonScalaObjectMapper[F <: JsonFactory]( + jsonFactory: F + ): JacksonScalaObjectMapperType = configureJacksonScalaObjectMapper(jsonFactory) + + private[this] final def configureJacksonScalaObjectMapper[ + F <: JsonFactory, + B <: TSFBuilder[F, B] + ]( + builder: TSFBuilder[F, B] + ): JacksonScalaObjectMapperType = configureJacksonScalaObjectMapper(builder.build()) + + private[jackson] final def configureJacksonScalaObjectMapper( + factory: JsonFactory + ): JacksonScalaObjectMapperType = + configureJacksonScalaObjectMapper( + new JacksonObjectMapper(factory) with JacksonScalaObjectMapper) + + private[jackson] final def configureJacksonScalaObjectMapper( + underlying: JacksonScalaObjectMapperType + ): JacksonScalaObjectMapperType = { + if (this.numbersAsStrings) { + underlying.enable(JsonWriteFeature.WRITE_NUMBERS_AS_STRINGS.mappedFeature()) + } + + this.defaultMapperConfiguration(underlying) + this.additionalMapperConfigurationFns.foreach(_(underlying)) + + underlying.setPropertyNamingStrategy(this.propertyNamingStrategy) + // Block use of a set of "unsafe" base types such as java.lang.Object + // to prevent exploitation of Remote Code Execution (RCE) vulnerability + // This line can be removed when this feature is enabled by default in Jackson 3 + underlying.enable(MapperFeature.BLOCK_UNSAFE_POLYMORPHIC_BASE_TYPES) + + this.jacksonModules.foreach(underlying.registerModule) + + underlying + } + } +} + +private[jackson] object ArrayElementsOnNewLinesPrettyPrinter extends DefaultPrettyPrinter { + _arrayIndenter = DefaultIndenter.SYSTEM_LINEFEED_INSTANCE + override def createInstance(): DefaultPrettyPrinter = new DefaultPrettyPrinter(this) +} + +/** + * A thin wrapper over a [[https://github.com/FasterXML/jackson-module-scala jackson-module-scala]] + * [[com.fasterxml.jackson.module.scala.ScalaObjectMapper]] + * + * @note this API is inspired by the [[https://github.com/codahale/jerkson Jerkson]] + * [[https://github.com/codahale/jerkson/blob/master/src/main/scala/com/codahale/jerkson/Parser.scala Parser]] + * + * @param underlying a configured [[JacksonScalaObjectMapperType]] + */ +class ScalaObjectMapper(val underlying: JacksonScalaObjectMapperType) { + assert(underlying != null, "Underlying ScalaObjectMapper cannot be null.") + + /** + * Constructed [[ObjectWriter]] that will serialize objects using specified pretty printer + * for indentation (or if null, no pretty printer). + */ + lazy val prettyObjectMapper: ObjectWriter = + underlying.writer(ArrayElementsOnNewLinesPrettyPrinter) + + /** Returns the currently configured [[PropertyNamingStrategy]] */ + def propertyNamingStrategy: PropertyNamingStrategy = + underlying.getPropertyNamingStrategy + + /** + * Factory method for constructing a [[com.fasterxml.jackson.databind.ObjectReader]] that will + * read or update instances of specified type, `T`. + * @tparam T the type for which to create a [[com.fasterxml.jackson.databind.ObjectReader]] + * @return the created [[com.fasterxml.jackson.databind.ObjectReader]]. + */ + def reader[T: Manifest]: ObjectReader = + underlying.readerFor[T] + + /** Read a value from a [[Buf]] into a type `T`. */ + def parse[T: Manifest](buf: Buf): T = + parse[T](Buf.ByteBuffer.Shared.extract(buf)) + + /** Read a value from a [[ByteBuffer]] into a type `T`. */ + def parse[T: Manifest](byteBuffer: ByteBuffer): T = { + val is = new ByteBufferBackedInputStream(byteBuffer) + underlying.readValue[T](is) + } + + /** Convert from a [[JsonNode]] into a type `T`. */ + def parse[T: Manifest](jsonNode: JsonNode): T = + convert[T](jsonNode) + + /** Read a value from an [[InputStream]] (caller is responsible for closing the stream) into a type `T`. */ + def parse[T: Manifest](inputStream: InputStream): T = + underlying.readValue[T](inputStream) + + /** Read a value from an Array[Byte] into a type `T`. */ + def parse[T: Manifest](bytes: Array[Byte]): T = + underlying.readValue[T](bytes) + + /** Read a value from a String into a type `T`. */ + def parse[T: Manifest](string: String): T = + underlying.readValue[T](string) + + /** Read a value from a [[JsonParser]] into a type `T`. */ + def parse[T: Manifest](jsonParser: JsonParser): T = + underlying.readValue[T](jsonParser) + + /** + * Convenience method for doing two-step conversion from given value, into an instance of given + * value type, [[JavaType]] if (but only if!) conversion is needed. If given value is already of + * requested type, the value is returned as is. + * + * This method is functionally similar to first serializing a given value into JSON, and then + * binding JSON data into a value of the given type, but should be more efficient since full + * serialization does not (need to) occur. However, the same converters (serializers, + * deserializers) will be used for data binding, meaning the same object mapper configuration + * works. + * + * Note: it is possible that in some cases behavior does differ from full + * serialize-then-deserialize cycle. It is not guaranteed, however, that the behavior is 100% + * the same -- the goal is just to allow efficient value conversions for structurally compatible + * Objects, according to standard Jackson configuration. + * + * Further note that this functionality is not designed to support "advanced" use cases, such as + * conversion of polymorphic values, or cases where Object Identity is used. + * + * @param from value from which to convert. + * @param toValueType type to be converted into. + * @return a new instance of type [[JavaType]] converted from the given [[Any]] type. + */ + def convert(from: Any, toValueType: JavaType): AnyRef = { + try { + underlying.convertValue(from, toValueType) + } catch { + case e: IllegalArgumentException if e.getCause != null => + throw e.getCause + } + } + + /** + * Convenience method for doing two-step conversion from a given value, into an instance of a given + * type, `T`. This is functionality equivalent to first serializing the given value into JSON, + * then binding JSON data into a value of the given type, but may be executed without fully + * serializing into JSON. The same converters (serializers, deserializers) will be used for + * data binding, meaning the same object mapper configuration works. + * + * Note: when a [[com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException]] + * is thrown inside of the the `ObjectMapper#convertValue` method, Jackson wraps the exception + * inside of an [[IllegalArgumentException]]. As such we unwrap to restore the original + * exception here. The wrapping occurs because the [[com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException]] + * is a sub-type of [[java.io.IOException]] (through extension of [[com.fasterxml.jackson.databind.JsonMappingException]] + * --> [[com.fasterxml.jackson.core.JsonProcessingException]] --> [[java.io.IOException]]. + * + * @param any the value to be converted. + * @tparam T the type to which to be converted. + * @return a new instance of type `T` converted from the given [[Any]] type. + * + * @see [[https://github.com/FasterXML/jackson-databind/blob/d70b9e65c5e089094ec7583fa6a38b2f484a96da/src/main/java/com/fasterxml/jackson/databind/ObjectMapper.java#L2167]] + */ + def convert[T: Manifest](any: Any): T = { + try { + underlying.convertValue[T](any) + } catch { + case e: IllegalArgumentException if e.getCause != null => + throw e.getCause + } + } + + /** + * Method that can be used to serialize any value as JSON output, using the output stream + * provided (using an encoding of [[com.fasterxml.jackson.core.JsonEncoding#UTF8]]. + * + * Note: this method does not close the underlying stream explicitly here; however, the + * [[com.fasterxml.jackson.core.JsonFactory]] this mapper uses may choose to close the stream + * depending on its settings (by default, it will try to close it when + * [[com.fasterxml.jackson.core.JsonGenerator]] constructed is closed). + * + * @param any the value to serialize. + * @param outputStream the [[OutputStream]] to which to serialize. + */ + def writeValue(any: Any, outputStream: OutputStream): Unit = + underlying.writeValue(outputStream, any) + + /** + * Method that can be used to serialize any value as a `Array[Byte]`. Functionally equivalent + * to calling [[JacksonObjectMapper#writeValue(Writer,Object)]] with a [[java.io.ByteArrayOutputStream]] and + * getting bytes, but more efficient. Encoding used will be UTF-8. + * + * @param any the value to serialize. + * @return the `Array[Byte]` representing the serialized value. + */ + def writeValueAsBytes(any: Any): Array[Byte] = + underlying.writeValueAsBytes(any) + + /** + * Method that can be used to serialize any value as a String. Functionally equivalent to calling + * [[JacksonObjectMapper#writeValue(Writer,Object)]] with a [[java.io.StringWriter]] + * and constructing String, but more efficient. + * + * @param any the value to serialize. + * @return the String representing the serialized value. + */ + def writeValueAsString(any: Any): String = + underlying.writeValueAsString(any) + + /** + * Method that can be used to serialize any value as a pretty printed String. Uses the + * [[prettyObjectMapper]] and calls [[writeValueAsString(any: Any)]] on the given value. + * + * @param any the value to serialize. + * @return the pretty printed String representing the serialized value. + * + * @see [[prettyObjectMapper]] + * @see [[writeValueAsString(any: Any)]] + */ + def writePrettyString(any: Any): String = any match { + case str: String => + val jsonNode = underlying.readValue[JsonNode](str) + prettyObjectMapper.writeValueAsString(jsonNode) + case _ => + prettyObjectMapper.writeValueAsString(any) + } + + /** + * Method that can be used to serialize any value as a [[Buf]]. Functionally equivalent + * to calling [[writeValueAsBytes(any: Any)]] and then wrapping the results in an + * "owned" [[Buf.ByteArray]]. + * + * @param any the value to serialize. + * @return the [[Buf.ByteArray.Owned]] representing the serialized value. + * + * @see [[writeValueAsBytes(any: Any)]] + * @see [[Buf.ByteArray.Owned]] + */ + def writeValueAsBuf(any: Any): Buf = + Buf.ByteArray.Owned(underlying.writeValueAsBytes(any)) + + /** + * Convenience method for doing the multi-step process of serializing a `Map[String, String]` + * to a [[Buf]]. + * + * @param stringMap the `Map[String, String]` to convert. + * @return the [[Buf.ByteArray.Owned]] representing the serialized value. + */ + // optimized + def writeStringMapAsBuf(stringMap: Map[String, String]): Buf = { + val os = new ByteArrayOutputStream() + + val jsonGenerator = underlying.getFactory.createGenerator(os) + try { + jsonGenerator.writeStartObject() + for ((key, value) <- stringMap) { + jsonGenerator.writeStringField(key, value) + } + jsonGenerator.writeEndObject() + jsonGenerator.flush() + + Buf.ByteArray.Owned(os.toByteArray) + } finally { + jsonGenerator.close() + } + } + + /** + * Method for registering a module that can extend functionality provided by this mapper; for + * example, by adding providers for custom serializers and deserializers. + * + * @note this mutates the [[underlying]] [[com.fasterxml.jackson.databind.ObjectMapper]] of + * this [[ScalaObjectMapper]]. + * + * @param module [[com.fasterxml.jackson.databind.Module]] to register. + */ + def registerModule(module: Module): JacksonObjectMapper = + underlying.registerModule(module) +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/YAML.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/YAML.scala new file mode 100644 index 0000000000..a15a4aae46 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/YAML.scala @@ -0,0 +1,65 @@ +package com.twitter.util.jackson + +import com.twitter.io.{Buf, ClasspathResource} +import com.twitter.util.Try +import java.io.{File, InputStream} + +/** + * Uses an instance of a [[ScalaObjectMapper]] configured to not perform any type of validation. + * Inspired by [[https://www.scala-lang.org/api/2.12.6/scala-parser-combinators/scala/util/parsing/json/JSON$.html]]. + * + * @note This is only intended for use from Scala (not Java). + * + * @see [[ScalaObjectMapper]] + */ +object YAML { + private final val Mapper: ScalaObjectMapper = + ScalaObjectMapper.builder.withNoValidation.yamlObjectMapper + + /** Simple utility to parse a YAML string into an Option[T] type. */ + def parse[T: Manifest](input: String): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** Simple utility to parse a YAML [[Buf]] into an Option[T] type. */ + def parse[T: Manifest](input: Buf): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** + * Simple utility to parse a YAML [[InputStream]] into an Option[T] type. + * @note the caller is responsible for managing the lifecycle of the given [[InputStream]]. + */ + def parse[T: Manifest](input: InputStream): Option[T] = + Try(Mapper.parse[T](input)).toOption + + /** + * Simple utility to load a YAML file and parse contents into an Option[T] type. + * + * @note the caller is responsible for managing the lifecycle of the given [[File]]. + */ + def parse[T: Manifest](f: File): Option[T] = + Try(Mapper.underlying.readValue[T](f)).toOption + + /** Simple utility to write a value as a YAML encoded String. */ + def write(any: Any): String = + Mapper.writeValueAsString(any) + + object Resource { + + /** + * Simple utility to load a YAML resource from the classpath and parse contents into an Option[T] type. + * + * @note `name` resolution to locate the resource is governed by [[java.lang.Class#getResourceAsStream]] + */ + def parse[T: Manifest](name: String): Option[T] = { + ClasspathResource.load(name) match { + case Some(inputStream) => + try { + YAML.parse[T](inputStream) + } finally { + inputStream.close() + } + case _ => None + } + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/BUILD b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/BUILD new file mode 100644 index 0000000000..898dffaa96 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/BUILD @@ -0,0 +1,33 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-caseclass", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/org/json4s:json4s-core", + "3rdparty/jvm/org/scala-lang/modules:scala-collection-compat", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-jackson-annotations/src/main/java/com/twitter/util/jackson/annotation", + "util/util-jackson/src/main/scala/com/fasterxml/jackson/databind", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-validator/src/main/scala/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation/conversions", + ], + exports = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + ], +) diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassBeanProperty.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassBeanProperty.scala new file mode 100644 index 0000000000..867af429b4 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassBeanProperty.scala @@ -0,0 +1,8 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.annotation.JacksonInject +import com.fasterxml.jackson.databind.BeanProperty + +private[twitter] case class CaseClassBeanProperty( + valueObj: Option[JacksonInject.Value], + property: BeanProperty) diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializer.scala new file mode 100644 index 0000000000..c80ebec7b3 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializer.scala @@ -0,0 +1,834 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.annotation.{JsonCreator, JsonIgnoreProperties} +import com.fasterxml.jackson.core.{ + JsonFactory, + JsonParseException, + JsonParser, + JsonProcessingException, + JsonToken +} +import com.fasterxml.jackson.databind._ +import com.fasterxml.jackson.databind.annotation.JsonNaming +import com.fasterxml.jackson.databind.exc.{ + InvalidDefinitionException, + InvalidFormatException, + MismatchedInputException, + UnrecognizedPropertyException +} +import com.fasterxml.jackson.databind.introspect._ +import com.fasterxml.jackson.databind.util.SimpleBeanPropertyDefinition +import com.twitter.util.WrappedValue +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.ValidationError +import com.twitter.util.jackson.caseclass.exceptions.{ + CaseClassFieldMappingException, + CaseClassMappingException, + InjectableValuesException, + _ +} +import com.twitter.util.logging.Logging +import com.twitter.util.reflect.{Annotations, Types => ReflectionTypes} +import com.twitter.util.validation.ScalaValidator +import com.twitter.util.validation.conversions.ConstraintViolationOps._ +import com.twitter.util.validation.conversions.PathOps._ +import com.twitter.util.validation.metadata.{ExecutableDescriptor, MethodDescriptor} +import jakarta.validation.{ConstraintViolation, Payload} +import java.lang.annotation.Annotation +import java.lang.reflect.{ + Constructor, + Executable, + InvocationTargetException, + Method, + Parameter, + ParameterizedType, + Type, + TypeVariable +} +import org.json4s.reflect.{ClassDescriptor, ConstructorDescriptor, Reflector, ScalaType} +import scala.collection.mutable +import scala.jdk.CollectionConverters._ +import scala.util.control.NonFatal + +private object CaseClassDeserializer { + // For reporting an InvalidDefinitionException + val EmptyJsonParser: JsonParser = new JsonFactory().createParser("") + + /* For supporting JsonCreator */ + case class CaseClassCreator( + executable: Executable, + propertyDefinitions: Array[PropertyDefinition], + executableValidationDescription: Option[ExecutableDescriptor], + executableValidationMethodDescriptions: Option[Array[MethodDescriptor]]) + + /* Holder for a fully specified JavaType with generics and a Jackson BeanPropertyDefinition */ + case class PropertyDefinition( + javaType: JavaType, + scalaType: ScalaType, + beanPropertyDefinition: BeanPropertyDefinition) + + /* Holder of constructor arg to ScalaType */ + case class ConstructorParam(name: String, scalaType: ScalaType) + + private[this] def applyPropertyNamingStrategy( + config: DeserializationConfig, + fieldName: String + ): String = + config.getPropertyNamingStrategy.nameForField(config, null, fieldName) + + // Validate `Constraint` annotations for all fields. + def executeFieldValidations( + validatorOption: Option[ScalaValidator], + executableDescriptorOption: Option[ExecutableDescriptor], + config: DeserializationConfig, + constructorValues: Array[Object], + fields: Array[CaseClassField] + ): Seq[CaseClassFieldMappingException] = (validatorOption, executableDescriptorOption) match { + case (Some(validator), Some(executableDescriptor)) => + val violations: Set[ConstraintViolation[Any]] = + validator.forExecutables + .validateExecutableParameters[Any]( + executableDescriptor, + constructorValues.asInstanceOf[Array[Any]], + fields.map(_.name) + ) + if (violations.nonEmpty) { + val results = mutable.ListBuffer[CaseClassFieldMappingException]() + val iterator = violations.iterator + while (iterator.hasNext) { + val violation: ConstraintViolation[Any] = iterator.next() + results.append( + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf( + applyPropertyNamingStrategy(config, violation.getPropertyPath.getLeafNode.toString) + ), + CaseClassFieldMappingException.Reason( + message = violation.getMessage, + detail = ValidationError( + violation, + CaseClassFieldMappingException.ValidationError.Field, + violation.getDynamicPayload(classOf[Payload]) + ) + ) + ) + ) + } + results.toSeq + } else Seq.empty[CaseClassFieldMappingException] + case _ => + Seq.empty[CaseClassFieldMappingException] + } + + // We do not want to "leak" method names, so we only want to return the leaf node + // if the ConstraintViolation PropertyPath size is greater than 1 element. + private[this] def getMethodValidationViolationPropertyPath( + config: DeserializationConfig, + violation: ConstraintViolation[_] + ): CaseClassFieldMappingException.PropertyPath = { + val propertyPath = violation.getPropertyPath + val iterator = propertyPath.iterator.asScala + if (iterator.hasNext) { + iterator.next() // move the iterator + if (iterator.hasNext) { // has another element, return the leaf + CaseClassFieldMappingException.PropertyPath.leaf( + applyPropertyNamingStrategy(config, propertyPath.getLeafNode.toString) + ) + } else { // does not have another element, return empty + CaseClassFieldMappingException.PropertyPath.Empty + } + } else { + CaseClassFieldMappingException.PropertyPath.Empty + } + } + + // Validate all methods annotated with `MethodValidation` defined in the deserialized case class. + // This is called after the case class is created. + def executeMethodValidations( + validatorOption: Option[ScalaValidator], + methodDescriptorsOption: Option[Array[MethodDescriptor]], + config: DeserializationConfig, + obj: Any + ): Seq[CaseClassFieldMappingException] = (validatorOption, methodDescriptorsOption) match { + case (Some(validator), Some(methodDescriptors)) => + try { + validator.forExecutables + .validateMethods[Any](methodDescriptors, obj).map { violation: ConstraintViolation[Any] => + CaseClassFieldMappingException( + getMethodValidationViolationPropertyPath(config, violation), + CaseClassFieldMappingException.Reason( + message = violation.getMessage, + detail = ValidationError( + violation, + CaseClassFieldMappingException.ValidationError.Method, + violation.getDynamicPayload(classOf[Payload])) + ) + ) + }.toSeq + } catch { + case NonFatal(e: InvocationTargetException) if e.getCause != null => + // propagate the InvocationTargetException + throw e.getCause + } + case _ => + Seq.empty[CaseClassFieldMappingException] + } +} + +/** + * Custom case class deserializer which overcomes limitations in jackson-scala-module. + * + * Our improvements: + * - Throws a [[JsonMappingException]] when non-`Option` fields are missing in the incoming JSON. + * - Does not allow for case class constructor fields to ever be constructed with a `null` reference. + * Any field which is marked to not be read from the incoming JSON must either be injected with a + * configured InjectableValues implementation or have a default value which can be supplied when + * the deserializer constructs the case class. Otherwise an error will be returned. + * - Use default values when fields are missing in the incoming JSON. + * - Properly deserialize a `Seq[Long]` (see https://github.com/FasterXML/jackson-module-scala/issues/62) + * - Support "wrapped values" using [[WrappedValue]]. + * - Support for field and method level validations. + * + * The following Jackson annotations which affect deserialization are not explicitly supported by + * this deserializer (note this list may not be exhaustive): + * - @JsonPOJOBuilder + * - @JsonAlias + * - @JsonSetter + * - @JsonAnySetter + * - @JsonTypeName + * - @JsonUnwrapped + * - @JsonManagedReference + * - @JsonBackReference + * - @JsonIdentityInfo + * - @JsonTypeIdResolver + * + * For many, this is because the behavior of these annotations is ambiguous when it comes to application on + * constructor arguments of a Scala case class during deserialization. + * + * @see [[https://github.com/FasterXML/jackson-annotations/wiki/Jackson-Annotations]] + * + * @note This class is inspired by Jerkson's CaseClassDeserializer which can be found here: + * [[https://github.com/codahale/jerkson/blob/master/src/main/scala/com/codahale/jerkson/deser/CaseClassDeserializer.scala]] + */ +/* exposed for testing */ +private[jackson] class CaseClassDeserializer( + javaType: JavaType, + config: DeserializationConfig, + beanDescription: BeanDescription, + validator: Option[ScalaValidator]) + extends JsonDeserializer[AnyRef] + with Logging { + import CaseClassDeserializer._ + + private[this] val clazz: Class[_] = javaType.getRawClass + // we explicitly do not read a mix-in for a primitive type + private[this] val mixinClazz: Option[Class[_]] = + Option(config.findMixInClassFor(clazz)) + .flatMap(m => if (m.isPrimitive) None else Some(m)) + private[this] val clazzDescriptor: ClassDescriptor = + Reflector.describe(clazz).asInstanceOf[ClassDescriptor] + private[this] val clazzAnnotations: Array[Annotation] = + mixinClazz.map(_.getAnnotations) match { + case Some(mixinAnnotations) => + clazz.getAnnotations ++ mixinAnnotations + case _ => + clazz.getAnnotations + } + + private[this] val caseClazzCreator: CaseClassCreator = { + val fromCompanion: Option[AnnotatedMethod] = + beanDescription.getFactoryMethods.asScala.find(_.hasAnnotation(classOf[JsonCreator])) + val fromClazz: Option[AnnotatedConstructor] = + beanDescription.getConstructors.asScala.find(_.hasAnnotation(classOf[JsonCreator])) + + fromCompanion match { + case Some(jsonCreatorAnnotatedMethod) => + CaseClassCreator( + jsonCreatorAnnotatedMethod.getAnnotated, + getBeanPropertyDefinitions( + jsonCreatorAnnotatedMethod.getAnnotated.getParameters, + jsonCreatorAnnotatedMethod, + fromCompanion = true), + validator.map(_.describeExecutable(jsonCreatorAnnotatedMethod.getAnnotated, mixinClazz)), + validator.map(_.describeMethods(clazz)) + ) + case _ => + fromClazz match { + case Some(jsonCreatorAnnotatedConstructor) => + CaseClassCreator( + jsonCreatorAnnotatedConstructor.getAnnotated, + getBeanPropertyDefinitions( + jsonCreatorAnnotatedConstructor.getAnnotated.getParameters, + jsonCreatorAnnotatedConstructor), + validator.map( + _.describeExecutable(jsonCreatorAnnotatedConstructor.getAnnotated, mixinClazz)), + validator.map(_.describeMethods(clazz)) + ) + case _ => + // try to use what Jackson thinks is the default -- however Jackson does not + // seem to correctly track an empty default constructor for case classes, nor + // multiple un-annotated and we have no way to pick a proper constructor so we bail + val constructor = Option(beanDescription.getClassInfo.getDefaultConstructor) match { + case Some(ctor) => ctor + case _ => + val constructors = beanDescription.getBeanClass.getConstructors + assert(constructors.size == 1, "Multiple case class constructors not supported") + annotateConstructor(constructors.head, clazzAnnotations) + } + CaseClassCreator( + constructor.getAnnotated, + getBeanPropertyDefinitions(constructor.getAnnotated.getParameters, constructor), + validator.map(_.describeExecutable(constructor.getAnnotated, mixinClazz)), + validator.map(_.describeMethods(clazz)) + ) + } + } + } + // nested class inside another class is not supported, e.g., we do not support + // use of creators for non-static inner classes, + assert(!beanDescription.isNonStaticInnerClass, "Non-static inner case classes are not supported.") + + // all class annotations -- including inherited annotations for fields + private[this] val allAnnotations: scala.collection.Map[String, Array[Annotation]] = { + // use the carried scalaType in order to find the appropriate constructor by the + // unresolved scalaTypes -- e.g., before we resolve generic types to their bound type + Annotations.findAnnotations( + clazz, + caseClazzCreator.propertyDefinitions + .map(_.scalaType.erasure) + ) + } + + // support for reading annotations from Jackson Mix-ins + // see: https://github.com/FasterXML/jackson-docs/wiki/JacksonMixInAnnotations + private[this] val allMixinAnnotations: scala.collection.Map[String, Array[Annotation]] = + mixinClazz match { + case Some(clazz) => + clazz.getDeclaredMethods.map(m => m.getName -> m.getAnnotations).toMap + case _ => + scala.collection.Map.empty[String, Array[Annotation]] + } + + // optimized lookup of bean property definition annotations + private[this] def getBeanPropertyDefinitionAnnotations( + beanPropertyDefinition: BeanPropertyDefinition + ): Array[Annotation] = { + if (beanPropertyDefinition.getPrimaryMember.getAllAnnotations.size() > 0) { + beanPropertyDefinition.getPrimaryMember.getAllAnnotations.annotations().asScala.toArray + } else Array.empty[Annotation] + } + + // Field name to list of parsed annotations. Jackson only tracks JacksonAnnotations + // in the BeanPropertyDefinition AnnotatedMembers and we want to track all class annotations by field. + // Annotations are keyed by parameter name because the logic collapses annotations from the + // inheritance hierarchy where the discriminator is member name. + // note: Prefer while loop over Scala for loop for better performance. The scala for loop + // performance is optimized in 2.13.0 if we enable scalac: https://github.com/scala/bug/issues/1338 + private[this] val fieldAnnotations: scala.collection.Map[String, Array[Annotation]] = { + val fieldBeanPropertyDefinitions: Array[BeanPropertyDefinition] = + caseClazzCreator.propertyDefinitions.map(_.beanPropertyDefinition) + val collectedFieldAnnotations = mutable.HashMap[String, Array[Annotation]]() + var index = 0 + while (index < fieldBeanPropertyDefinitions.length) { + val fieldBeanPropertyDefinition = fieldBeanPropertyDefinitions(index) + val fieldName = fieldBeanPropertyDefinition.getInternalName + // in many cases we will have the field annotations in the `allAnnotations` Map which + // scans for annotations from the class definition being deserialized, however in some cases + // we are dealing with an annotated field from a static or secondary constructor + // and thus the annotations may not exist in the `allAnnotations` Map so we default to + // any carried bean property definition annotations which includes annotation information + // for any static or secondary constructor. + val annotations = + allAnnotations.getOrElse( + fieldName, + getBeanPropertyDefinitionAnnotations(fieldBeanPropertyDefinition) + ) ++ allMixinAnnotations.getOrElse(fieldName, Array.empty) + if (annotations.nonEmpty) collectedFieldAnnotations.put(fieldName, annotations) + index += 1 + } + + collectedFieldAnnotations + } + + /* exposed for testing */ + private[jackson] val fields: Array[CaseClassField] = + CaseClassField.createFields( + clazz, + caseClazzCreator.executable, + clazzDescriptor, + caseClazzCreator.propertyDefinitions, + fieldAnnotations, + propertyNamingStrategy + ) + + private[this] lazy val numConstructorArgs = fields.length + private[this] lazy val isWrapperClass = classOf[WrappedValue[_]].isAssignableFrom(clazz) + private[this] lazy val firstFieldName = fields.head.name + + /* Public */ + + override def isCachable: Boolean = true + + override def deserialize(jsonParser: JsonParser, context: DeserializationContext): Object = { + if (isWrapperClass) + deserializeWrapperClass(jsonParser, context) + else + deserializeNonWrapperClass(jsonParser, context) + } + + /* Private */ + + private[this] def deserializeWrapperClass( + jsonParser: JsonParser, + context: DeserializationContext + ): Object = { + if (jsonParser.getCurrentToken.isStructStart) { + try { + context.handleUnexpectedToken( + clazz, + jsonParser.currentToken(), + jsonParser, + "Unable to deserialize wrapped value from a json object" + ) + } catch { + case NonFatal(e) => + // wrap in a JsonMappingException + throw JsonMappingException.from(jsonParser, e.getMessage) + } + } + + val jsonNode = context.getNodeFactory.objectNode() + jsonNode.put(firstFieldName, jsonParser.getText) + deserialize(jsonParser, context, jsonNode) + } + + private[this] def deserializeNonWrapperClass( + jsonParser: JsonParser, + context: DeserializationContext + ): Object = { + incrementParserToFirstField(jsonParser, context) + val jsonNode = jsonParser.readValueAsTree[JsonNode] + deserialize(jsonParser, context, jsonNode) + } + + private[this] def deserialize( + jsonParser: JsonParser, + context: DeserializationContext, + jsonNode: JsonNode + ): Object = { + val jsonFieldNames: Seq[String] = jsonNode.fieldNames().asScala.toSeq + val caseClassFieldNames: Seq[String] = fields.map(_.name) + + handleUnknownFields(jsonParser, context, jsonFieldNames, caseClassFieldNames) + val (values, parseErrors) = parseConstructorValues(jsonParser, context, jsonNode) + + // run field validations + val validationErrors = executeFieldValidations( + validator, + caseClazzCreator.executableValidationDescription, + config, + values, + fields) + + val errors = parseErrors ++ validationErrors + if (errors.nonEmpty) throw CaseClassMappingException(errors.toSet) + + newInstance(values) + } + + /** Return all "unknown" properties sent in the incoming JSON */ + private[this] def unknownProperties( + context: DeserializationContext, + jsonFieldNames: Seq[String], + caseClassFieldNames: Seq[String] + ): Seq[String] = { + // if there is a JsonIgnoreProperties annotation on the class, it should prevail + val nonIgnoredFields: Seq[String] = + Annotations.findAnnotation[JsonIgnoreProperties](clazzAnnotations) match { + case Some(annotation) if !annotation.ignoreUnknown() => + // has a JsonIgnoreProperties annotation and is configured to NOT ignore unknown properties + val annotationIgnoredFields: Seq[String] = annotation.value() + // only non-ignored json fields should be considered + jsonFieldNames.diff(annotationIgnoredFields) + case None if context.isEnabled(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES) => + // no annotation but feature is configured, thus all json fields should be considered + jsonFieldNames + case _ => + Seq.empty[String] // every field is ignorable + } + + // if we have any non ignored field, return the difference + if (nonIgnoredFields.nonEmpty) { + nonIgnoredFields.diff(caseClassFieldNames) + } else Seq.empty[String] + } + + /** Throws a [[CaseClassMappingException]] with a set of [[CaseClassFieldMappingException]] per unknown field */ + private[this] def handleUnknownFields( + jsonParser: JsonParser, + context: DeserializationContext, + jsonFieldNames: Seq[String], + caseClassFieldNames: Seq[String] + ): Unit = { + // handles more or less incoming fields in the JSON than are defined in the case class + val unknownFields: Seq[String] = unknownProperties(context, jsonFieldNames, caseClassFieldNames) + if (unknownFields.nonEmpty) { + val errors = unknownFields.map { field => + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.Empty, + CaseClassFieldMappingException.Reason( + message = UnrecognizedPropertyException + .from( + jsonParser, + clazz, + field, + caseClassFieldNames.map(_.asInstanceOf[Object]).asJavaCollection + ).getMessage + ) + ) + } + throw CaseClassMappingException(errors.toSet) + } + } + + private[this] def incrementParserToFirstField( + jsonParser: JsonParser, + context: DeserializationContext + ): Unit = { + if (jsonParser.getCurrentToken == JsonToken.START_OBJECT) { + jsonParser.nextToken() + } + if (jsonParser.getCurrentToken != JsonToken.FIELD_NAME && + jsonParser.getCurrentToken != JsonToken.END_OBJECT) { + try { + context.handleUnexpectedToken(clazz, jsonParser) + } catch { + case NonFatal(e) => + e match { + case j: JsonProcessingException => + // don't include source info since it's often blank. + j.clearLocation() + throw new JsonParseException(jsonParser, j.getMessage) + case _ => + throw new JsonParseException(jsonParser, e.getMessage) + } + } + } + } + + // note: Prefer while loop over Scala for loop for better performance. The scala for loop + // performance is optimized in 2.13.0 if we enable scalac: https://github.com/scala/bug/issues/1338 + private[this] def parseConstructorValues( + jsonParser: JsonParser, + context: DeserializationContext, + jsonNode: JsonNode + ): (Array[Object], Seq[CaseClassFieldMappingException]) = { + /* Mutable Fields */ + var constructorValuesIdx = 0 + val constructorValues = new Array[Object](numConstructorArgs) + val errors = mutable.ListBuffer[CaseClassFieldMappingException]() + + while (constructorValuesIdx < numConstructorArgs) { + val field = fields(constructorValuesIdx) + try { + val value = field.parse(context, jsonParser.getCodec, jsonNode) + constructorValues(constructorValuesIdx) = value //mutation + } catch { + case e: CaseClassFieldMappingException => + if (e.path == null) { + // fill in missing path detail + addException( + field, + e.withPropertyPath(CaseClassFieldMappingException.PropertyPath.leaf(field.name)), + constructorValues, + constructorValuesIdx, + errors + ) + } else { + addException( + field, + e, + constructorValues, + constructorValuesIdx, + errors + ) + } + case e: InvalidFormatException => + addException( + field, + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf(field.name), + CaseClassFieldMappingException.Reason( + s"'${e.getValue.toString}' is not a " + + s"valid ${Types.wrapperType(e.getTargetType).getSimpleName}${validValuesString(e)}", + CaseClassFieldMappingException.JsonProcessingError(e) + ) + ), + constructorValues, + constructorValuesIdx, + errors + ) + case e: MismatchedInputException => + addException( + field, + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf(field.name), + CaseClassFieldMappingException.Reason( + s"'${jsonNode.asText("")}' is not a " + + s"valid ${Types.wrapperType(e.getTargetType).getSimpleName}${validValuesString(e)}", + CaseClassFieldMappingException.JsonProcessingError(e) + ) + ), + constructorValues, + constructorValuesIdx, + errors + ) + case e: CaseClassMappingException => + constructorValues(constructorValuesIdx) = field.missingValue //mutation + errors ++= e.errors.map(_.scoped(field.name)) + case e: JsonProcessingException => + // don't include source info since it's often blank. Consider adding e.getCause.getMessage + e.clearLocation() + addException( + field, + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf(field.name), + CaseClassFieldMappingException.Reason( + e.errorMessage, + CaseClassFieldMappingException.JsonProcessingError(e) + ) + ), + constructorValues, + constructorValuesIdx, + errors + ) + case e: java.util.NoSuchElementException if isScalaEnumerationType(field) => + // Scala enumeration mapping issue + addException( + field, + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf(field.name), + CaseClassFieldMappingException.Reason(e.getMessage) + ), + constructorValues, + constructorValuesIdx, + errors + ) + case e: InjectableValuesException => + // we rethrow, to prevent leaking internal injection detail in the "errors" array + throw e + case NonFatal(e) => + error("Unexpected exception parsing field: " + field, e) + throw e + } + constructorValuesIdx += 1 + } + + (constructorValues, errors.toSeq) + } + + private[this] def isScalaEnumerationType(field: CaseClassField): Boolean = { + // scala.Enumerations are challenging for class type comparison and thus we simply check + // if the class name "starts with" scala.Enumeration which should indicate if a class + // is a scala.Enumeration type. The added benefit is that this should work even if the + // classes come from different classloaders. + field.javaType.getRawClass.getName.startsWith(classOf[scala.Enumeration].getName) || + (field.isOption && field.javaType + .containedType(0).getRawClass.getName.startsWith(classOf[scala.Enumeration].getName)) + } + + /** Add the given exception to the given array buffer of errors while also adding a missing value field to the given array */ + private[this] def addException( + field: CaseClassField, + e: CaseClassFieldMappingException, + array: Array[Object], + idx: Int, + errors: mutable.ListBuffer[CaseClassFieldMappingException] + ): Unit = { + array(idx) = field.missingValue //mutation + errors += e //mutation + } + + private[this] def validValuesString(e: MismatchedInputException): String = { + if (e.getTargetType != null && e.getTargetType.isEnum) + " with valid values: " + e.getTargetType.getEnumConstants.mkString(", ") + else + "" + } + + private[this] def newInstance(constructorValues: Array[Object]): Object = { + val obj = + try { + instantiate(constructorValues) + } catch { + case e @ (_: InvocationTargetException | _: ExceptionInInitializerError) => + if (e.getCause == null) + throw e // propagate the underlying cause of the failed instantiation if available + else throw e.getCause + } + val methodValidationErrors = + executeMethodValidations( + validator, + caseClazzCreator.executableValidationMethodDescriptions, + config, + obj + ) + if (methodValidationErrors.nonEmpty) + throw CaseClassMappingException(methodValidationErrors.toSet) + + obj + } + + private[this] def instantiate( + constructorValues: Array[Object] + ): Object = { + caseClazzCreator.executable match { + case method: Method => + // if the creator is of type Method, we assume the need to invoke the companion object + method.invoke(clazzDescriptor.companion.get.instance, constructorValues: _*) + case constructor: Constructor[_] => + // otherwise simply invoke the constructor + constructor.newInstance(constructorValues: _*).asInstanceOf[Object] + } + } + + private[this] def propertyNamingStrategy: PropertyNamingStrategy = { + Annotations.findAnnotation[JsonNaming](clazzAnnotations) match { + case Some(jsonNaming) => + jsonNaming.value().getDeclaredConstructor().newInstance() + case _ => + config.getPropertyNamingStrategy + } + } + + private[this] def annotateConstructor( + constructor: Constructor[_], + annotations: Seq[Annotation] + ): AnnotatedConstructor = { + val paramAnnotationMaps: Array[AnnotationMap] = { + constructor.getParameterAnnotations.map { parameterAnnotations => + val parameterAnnotationMap = new AnnotationMap() + parameterAnnotations.map(parameterAnnotationMap.add) + parameterAnnotationMap + } + } + + val annotationMap = new AnnotationMap() + annotations.map(annotationMap.add) + new AnnotatedConstructor( + new TypeResolutionContext.Basic(config.getTypeFactory, javaType.getBindings), + constructor, + annotationMap, + paramAnnotationMaps + ) + } + + /* in order to deal with parameterized types we create a JavaType here and carry it */ + private[this] def getBeanPropertyDefinitions( + parameters: Array[Parameter], + annotatedWithParams: AnnotatedWithParams, + fromCompanion: Boolean = false + ): Array[PropertyDefinition] = { + // need to find the scala description which carries the full type information + val constructorParamDescriptors = + findConstructorDescriptor(parameters) match { + case Some(constructorDescriptor) => + constructorDescriptor.params + case _ => + throw InvalidDefinitionException.from( + CaseClassDeserializer.EmptyJsonParser, + s"Unable to locate suitable constructor for class: ${clazz.getName}", + javaType) + } + + for ((parameter, index) <- parameters.zipWithIndex) yield { + val constructorParamDescriptor = constructorParamDescriptors(index) + val scalaType = constructorParamDescriptor.argType + + val parameterJavaType = + if (!javaType.getBindings.isEmpty && + shouldFullyDefineParameterizedType(scalaType, parameter)) { + // what types are bound to the generic case class parameters + val boundTypeParameters: Array[JavaType] = + ReflectionTypes + .parameterizedTypeNames(parameter.getParameterizedType) + .map(javaType.getBindings.findBoundType) + Types + .javaType( + config.getTypeFactory, + scalaType, + boundTypeParameters + ) + } else { + Types.javaType(config.getTypeFactory, scalaType) + } + + val annotatedParameter = + newAnnotatedParameter( + typeResolutionContext = new TypeResolutionContext.Basic( + config.getTypeFactory, + javaType.getBindings + ), // use the TypeBindings from the top-level JavaType, not the parameter JavaType + owner = annotatedWithParams, + annotations = annotatedWithParams.getParameterAnnotations(index), + javaType = parameterJavaType, + index = index + ) + + PropertyDefinition( + parameterJavaType, + scalaType, + SimpleBeanPropertyDefinition + .construct( + config, + annotatedParameter, + new PropertyName(constructorParamDescriptor.name) + ) + ) + } + } + + // note: Prefer while loop over Scala for loop for better performance. The scala for loop + // performance is optimized in 2.13.0 if we enable scalac: https://github.com/scala/bug/issues/1338 + private[this] def findConstructorDescriptor( + parameters: Array[Parameter] + ): Option[ConstructorDescriptor] = { + val constructors = clazzDescriptor.constructors + var index = 0 + while (index < constructors.length) { + val constructorDescriptor = constructors(index) + val params = constructorDescriptor.params + if (params.length == parameters.length) { + // description has the same number of parameters we're looking for, check each type in order + val checkedParams = params.map { param => + Types + .wrapperType(param.argType.erasure) + .isAssignableFrom(Types.wrapperType(parameters(param.argIndex).getType)) + } + if (checkedParams.forall(_ == true)) return Some(constructorDescriptor) + } + index += 1 + } + None + } + + // if we need to attempt to fully specify the JavaType because it is generically types + private[this] def shouldFullyDefineParameterizedType( + scalaType: ScalaType, + parameter: Parameter + ): Boolean = { + // only need to fully specify if the type is parameterized and it has more than one type arg + // or its typeArg is also parameterized. + def isParameterized(reflectionType: Type): Boolean = reflectionType match { + case _: ParameterizedType | _: TypeVariable[_] => + true + case _ => + false + } + + val parameterizedType = parameter.getParameterizedType + !scalaType.isPrimitive && + (scalaType.typeArgs.isEmpty || + scalaType.typeArgs.head.erasure.isAssignableFrom(classOf[Object])) && + parameterizedType != parameter.getType && + isParameterized(parameterizedType) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializerResolver.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializerResolver.scala new file mode 100644 index 0000000000..72d8295968 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassDeserializerResolver.scala @@ -0,0 +1,22 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.databind.deser.Deserializers +import com.fasterxml.jackson.databind.{BeanDescription, DeserializationConfig, JavaType} +import com.twitter.util.reflect.{Types => ReflectionTypes} +import com.twitter.util.validation.ScalaValidator + +private[jackson] class CaseClassDeserializerResolver( + validator: Option[ScalaValidator]) + extends Deserializers.Base { + + override def findBeanDeserializer( + javaType: JavaType, + deserializationConfig: DeserializationConfig, + beanDescription: BeanDescription + ): CaseClassDeserializer = { + if (ReflectionTypes.isCaseClass(javaType.getRawClass)) + new CaseClassDeserializer(javaType, deserializationConfig, beanDescription, validator) + else + null + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassField.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassField.scala new file mode 100644 index 0000000000..820827f5d3 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassField.scala @@ -0,0 +1,459 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.annotation.{ + JsonIgnore, + JsonIgnoreType, + JsonProperty, + JsonTypeInfo, + JsonView +} +import com.fasterxml.jackson.core.{JsonParser, ObjectCodec} +import com.fasterxml.jackson.databind.annotation.JsonDeserialize +import com.fasterxml.jackson.databind.introspect.BeanPropertyDefinition +import com.fasterxml.jackson.databind.node.TreeTraversingParser +import com.fasterxml.jackson.databind.util.ClassUtil +import com.fasterxml.jackson.databind.{BeanProperty => JacksonBeanProperty, _} +import com.twitter.conversions.StringOps._ +import com.twitter.util.jackson.annotation.InjectableValue +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException +import com.twitter.util.logging.Logging +import com.twitter.util.reflect.Annotations +import java.lang.annotation.Annotation +import java.lang.reflect.Executable +import org.json4s.reflect.ClassDescriptor +import scala.annotation.tailrec +import scala.reflect.NameTransformer + +/* Keeps a mapping of field "type" as a String and the field name */ +private case class FieldInfo(`type`: String, fieldName: String) + +object CaseClassField { + + // optimized + private[jackson] def createFields( + clazz: Class[_], + constructor: Executable, + clazzDescriptor: ClassDescriptor, + propertyDefinitions: Array[CaseClassDeserializer.PropertyDefinition], + fieldAnnotations: scala.collection.Map[String, Array[Annotation]], + namingStrategy: PropertyNamingStrategy + ): Array[CaseClassField] = { + /* CaseClassFields MUST be returned in constructor/method parameter invocation order */ + // we could choose to use the size of the property definitions here, but since they are + // computed there is the possibility it is incorrect and thus we want to ensure we have + // as many definitions as there are parameters defined for the given constructor. + val parameters = constructor.getParameters + val fields = new Array[CaseClassField](parameters.length) + + var index = 0 + while (index < parameters.length) { + val propertyDefinition = propertyDefinitions(index) + // we look up annotations by the field name as that is how they are keyed + val annotations: Array[Annotation] = + fieldAnnotations.get(propertyDefinition.beanPropertyDefinition.getName) match { + case Some(v) => v + case _ => Array.empty[Annotation] + } + val name = jsonNameForField( + namingStrategy, + annotations, + propertyDefinition.beanPropertyDefinition.getName) + + fields(index) = CaseClassField( + name = name, + index = index, + javaType = propertyDefinition.javaType, + parentClass = clazz, + annotations = annotations, + beanPropertyDefinition = propertyDefinition.beanPropertyDefinition, + defaultFn = defaultMethod(clazzDescriptor, index) + ) + + index += 1 + } + + fields + } + + private[this] def jsonNameForField( + namingStrategy: PropertyNamingStrategy, + annotations: Array[Annotation], + name: String + ): String = { + Annotations.findAnnotation[JsonProperty](annotations) match { + case Some(jsonProperty) if jsonProperty.value.nonEmpty => + jsonProperty.value + case _ => + val decodedName = NameTransformer.decode(name) // decode unicode escaped field names + // apply json naming strategy (e.g. snake_case) + namingStrategy.nameForField( + /* config = */ null, + /* field = */ null, + /* defaultName = */ decodedName) + } + } +} + +/* exposed for testing */ +private[jackson] case class CaseClassField private ( + name: String, + index: Int, + javaType: JavaType, + parentClass: Class[_], + annotations: Array[Annotation], + beanPropertyDefinition: BeanPropertyDefinition, + defaultFn: Option[() => Object]) + extends Logging { + + lazy val missingValue: AnyRef = { + if (javaType.isPrimitive) + ClassUtil.defaultValue(javaType.getRawClass) + else + null + } + + private[caseclass] val isOption = javaType.hasRawClass(classOf[Option[_]]) + private[this] val isString = javaType.hasRawClass(classOf[String]) + + /** Lazy as we may not have a contained type */ + private[this] lazy val firstTypeParam: JavaType = javaType.containedType(0) + + /** Annotations defined on the case class definition */ + private[this] lazy val clazzAnnotations: Array[Annotation] = + if (isOption) { + javaType.containedType(0).getRawClass.getAnnotations + } else { + javaType.getRawClass.getAnnotations + } + + /** If the field has a [[JsonIgnore(true)]] annotation or the type has been annotated with [[JsonIgnoreType(true)]]*/ + private[this] lazy val isIgnored: Boolean = + Annotations + .findAnnotation[JsonIgnore](annotations).exists(_.value) || + Annotations + .findAnnotation[JsonIgnoreType](clazzAnnotations).exists(_.value) + + /** If the field is annotated directly with [[JsonDeserialize]] or if the class is annotated similarly */ + private[this] lazy val jsonDeserializer: Option[JsonDeserialize] = + Annotations + .findAnnotation[JsonDeserialize](annotations) + .orElse(Annotations.findAnnotation[JsonDeserialize](clazzAnnotations)) + + /* Public */ + + /** + * Parse the field from a JsonNode representing a JSON object. NOTE: We would normally return a + * `Try[Object]`, but instead we use exceptions to optimize the non-failure case. + * + * @param context DeserializationContext for deserialization + * @param codec Codec for field + * @param objectJsonNode The JSON object + * + * @return The parsed object for this field + * @note Option fields default to None even if no default is specified + * + * @throws CaseClassFieldMappingException with reason for the parsing error + */ + def parse( + context: DeserializationContext, + codec: ObjectCodec, + objectJsonNode: JsonNode + ): Object = { + val injectableValuesOpt: Option[InjectableValues] = + new DeserializationContextAccessor(context).getInjectableValues + val forProperty: CaseClassBeanProperty = beanProperty(context) + + injectableValuesOpt match { + case Some(_) => + Option( + context.findInjectableValue( + /* valueId = */ forProperty.valueObj.map(_.getId).orNull, + /* forProperty = */ forProperty.property, + /* beanInstance = */ null)) match { + case Some(value) => + value + case _ => + findValue(context, codec, objectJsonNode, forProperty.property) + } + case _ => + findValue(context, codec, objectJsonNode, forProperty.property) + } + } + + /* Private */ + + private[this] def findValue( + context: DeserializationContext, + codec: ObjectCodec, + objectJsonNode: JsonNode, + beanProperty: JacksonBeanProperty + ): Object = { + // current active view (if any) + val activeJsonView: Option[Class[_]] = Option(context.getActiveView) + // current @JsonView annotation (if any) + val fieldJsonViews: Option[Seq[Class[_]]] = + Annotations.findAnnotation[JsonView](annotations).map(_.value.toSeq) + + // context has an active view *and* the field is annotated + if (activeJsonView.isDefined && fieldJsonViews.isDefined) { + if (activeJsonView.exists(fieldJsonViews.get.contains(_))) { + // active view is in the list of views from the annotation + parse(context, codec, objectJsonNode, beanProperty) + } else defaultValueOrException(isIgnored) + } else { + // no active view proceed as normal + parse(context, codec, objectJsonNode, beanProperty) + } + } + + private[this] def beanProperty( + context: DeserializationContext, + optionalJavaType: Option[JavaType] = None + ): CaseClassBeanProperty = + newBeanProperty( + context = context, + javaType = javaType, + optionalJavaType = optionalJavaType, + annotatedParameter = beanPropertyDefinition.getConstructorParameter, + annotations = annotations, + name = name, + index = index + ) + + private[this] def parse( + context: DeserializationContext, + codec: ObjectCodec, + objectJsonNode: JsonNode, + forProperty: JacksonBeanProperty + ): Object = { + val fieldJsonNode = objectJsonNode.get(name) + if (!isIgnored && fieldJsonNode != null) { + // not ignored and a value was passed in the JSON + resolveWithDeserializerAnnotation(context, codec, fieldJsonNode, forProperty) match { + case Some(obj) => obj + case _ => // wasn't handled by another resolved deserializer + parseFieldValue(context, codec, fieldJsonNode, forProperty, None) + } + } else defaultValueOrException(isIgnored) + } + + //optimized + private[this] def resolveWithDeserializerAnnotation( + context: DeserializationContext, + fieldCodec: ObjectCodec, + fieldJsonNode: JsonNode, + forProperty: JacksonBeanProperty + ): Option[Object] = jsonDeserializer match { + case Some(annotation: JsonDeserialize) + if isNotAssignableFrom(annotation.using, classOf[JsonDeserializer.None]) => + // specifies a deserializer to use. we don't want to run any processing of the field + // node value here as the field is specified to be deserialized by some other deserializer + Option( + context + .deserializerInstance(beanPropertyDefinition.getPrimaryMember, annotation.using)) match { + case Some(deserializer) => + val treeTraversingParser = new TreeTraversingParser(fieldJsonNode, fieldCodec) + try { + // advance the parser to the next token for deserialization + treeTraversingParser.nextToken + if (isOption) { + Some( + Option( + deserializer + .deserialize(treeTraversingParser, context))) + } else { + Some(deserializer.deserialize(treeTraversingParser, context)) + } + } finally { + treeTraversingParser.close() + } + case _ => + Some( + context.handleInstantiationProblem( + javaType.getRawClass, + annotation.using.toString, + JsonMappingException.from( + context, + "Unable to locate/create deserializer specified by: " + + s"${annotation.getClass.getName}(using = ${annotation.using()})") + ) + ) + } + case Some(annotation: JsonDeserialize) + if isNotAssignableFrom(annotation.contentAs, classOf[java.lang.Void]) => + // there is a @JsonDeserialize annotation but it merely states to deserialize as a specific type + Some( + parseFieldValue( + context, + fieldCodec, + fieldJsonNode, + forProperty, + Some(annotation.contentAs) + ) + ) + case _ => None + } + + private[this] def parseFieldValue( + context: DeserializationContext, + fieldCodec: ObjectCodec, + fieldJsonNode: JsonNode, + forProperty: JacksonBeanProperty, + subTypeClazz: Option[Class[_]] + ): Object = { + if (fieldJsonNode.isNull) { + // the passed JSON value is a 'null' value + defaultValueOrException(isIgnored) + } else if (isString) { + // if this is a String type, do not try to use a JsonParser and simply return the node as text + if (fieldJsonNode.isValueNode) fieldJsonNode.asText() + else fieldJsonNode.toString + } else { + val treeTraversingParser = new TreeTraversingParser(fieldJsonNode, fieldCodec) + try { + // advance the parser to the next token for deserialization + treeTraversingParser.nextToken + if (isOption) { + Option( + parseFieldValue( + context, + fieldCodec, + treeTraversingParser, + beanProperty(context, Some(firstTypeParam)).property, + subTypeClazz) + ) + } else { + assertNotNull( + context, + fieldJsonNode, + parseFieldValue(context, fieldCodec, treeTraversingParser, forProperty, subTypeClazz) + ) + } + } finally { + treeTraversingParser.close() + } + } + } + + private[this] def parseFieldValue( + context: DeserializationContext, + fieldCodec: ObjectCodec, + jsonParser: JsonParser, + forProperty: JacksonBeanProperty, + subTypeClazz: Option[Class[_]] + ): Object = { + val resolvedType = resolveSubType(context, forProperty.getType, subTypeClazz) + Annotations.findAnnotation[JsonTypeInfo](clazzAnnotations) match { + case Some(_) => + // for polymorphic types we cannot contextualize + // thus we go back to the field codec to read + fieldCodec.readValue(jsonParser, resolvedType) + case _ if resolvedType.isContainerType => + // nor container types -- trying to contextualize on a container type leads to really bad performance + // thus we go back to the field codec to read + fieldCodec.readValue(jsonParser, resolvedType) + case _ => + // contextualization for all others + context.readPropertyValue(jsonParser, forProperty, resolvedType) + } + } + + //optimized + private[this] def assertNotNull( + context: DeserializationContext, + field: JsonNode, + value: Object + ): Object = { + value match { + case null => + throw JsonMappingException.from(context, "error parsing '" + field.asText + "'") + case traversable: Traversable[_] => + assertNotNull(context, traversable) + case array: Array[_] => + assertNotNull(context, array) + case _ => // no-op + } + value + } + + private[this] def assertNotNull( + context: DeserializationContext, + traversable: Traversable[_] + ): Unit = { + if (traversable.exists(_ == null)) { + throw JsonMappingException.from( + context, + "Literal null values are not allowed as json array elements." + ) + } + } + + /** Return the default value of the field */ + private[this] def defaultValue: Option[Object] = { + if (defaultFn.isDefined) + defaultFn.map(_()) + else if (isOption) + Some(None) + else + None + } + + /** Return the [[defaultValue]] or else throw a required field exception */ + private[this] def defaultValueOrException(ignoredField: Boolean): Object = { + defaultValue match { + case Some(value) => value + case _ => throwRequiredFieldException(ignoredField) + } + } + + /** Throw a required field exception */ + private[this] def throwRequiredFieldException(ignorable: Boolean): Nothing = { + val (fieldInfoAttributeType, fieldInfoAttributeName) = getFieldInfo(name, annotations) + + val message = + if (ignorable) s"ignored $fieldInfoAttributeType has no default value specified" + else s"$fieldInfoAttributeType is required" + throw CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf(fieldInfoAttributeName), + CaseClassFieldMappingException.Reason( + message, + CaseClassFieldMappingException.RequiredFieldMissing + ) + ) + } + + @tailrec + private[this] def getFieldInfo( + fieldName: String, + annotations: Seq[Annotation] + ): (String, String) = { + if (annotations.isEmpty) { + ("field", fieldName) + } else { + extractFieldInfo(fieldName, annotations.head) match { + case Some(found) => found + case _ => getFieldInfo(fieldName, annotations.tail) + } + } + } + + // Annotations annotated with `@InjectableValue` represent values that are injected + // into a constructed case class from somewhere other than the incoming JSON, + // e.g., with Jackson InjectableValues. When we are parsing the structure of the + // case class, we want to capture information about these annotations for proper + // error reporting on the fields annotated. The annotation name is used instead + // of the generic classifier of "field", and the annotation#value() is used + // instead of the case class field name being marshalled into. + private[this] def extractFieldInfo( + fieldName: String, + annotation: Annotation + ): Option[(String, String)] = { + if (Annotations.isAnnotationPresent[InjectableValue](annotation)) { + val name = Annotations.getValueIfAnnotatedWith[InjectableValue](annotation) match { + case Some(value) if value != null && value.nonEmpty => value + case _ => fieldName + } + Some((annotation.annotationType.getSimpleName.toCamelCase, name)) + } else None + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassJacksonModule.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassJacksonModule.scala new file mode 100644 index 0000000000..ae77e1da89 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/CaseClassJacksonModule.scala @@ -0,0 +1,14 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.module.scala._ +import com.twitter.util.validation.ScalaValidator + +private[jackson] class CaseClassJacksonModule( + validator: Option[ScalaValidator]) + extends JacksonModule { + override def getModuleName: String = this.getClass.getName + + this += { + _.addDeserializers(new CaseClassDeserializerResolver(validator)) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/SerdeLogging.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/SerdeLogging.scala new file mode 100644 index 0000000000..fb088d8331 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/SerdeLogging.scala @@ -0,0 +1,24 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties + +/** + * "Serde" (technically "SerDe") is a portmanteau of "Serialization and Deserialization" + * + * Mix this trait into a case class to get helpful logging methods. + * This trait adds a `@JsonIgnoreProperties` for fields which are + * defined in the [[com.twitter.util.logging.Logging]] trait so that + * they are not included in JSON serde operations. + * + */ +@JsonIgnoreProperties( + Array( + "logger_name", + "trace_enabled", + "debug_enabled", + "error_enabled", + "info_enabled", + "warn_enabled" + ) +) +trait SerdeLogging extends com.twitter.util.logging.Logging diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/Types.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/Types.scala new file mode 100644 index 0000000000..907c1984f2 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/Types.scala @@ -0,0 +1,144 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.databind.JavaType +import com.fasterxml.jackson.databind.`type`.{ArrayType, TypeBindings, TypeFactory} +import org.json4s.reflect.ScalaType + +private[twitter] object Types { + + /** + * Primitive class to Java wrapper (boxed) primitive class. + * Similar to [[com.fasterxml.jackson.databind.util.ClassUtil.wrapperType(Class primitiveType)]] + * but accepts non primitives as a pass-thru. + */ + def wrapperType(clazz: Class[_]): Class[_] = { + if (clazz == null) + classOf[Null] + else + clazz match { + case java.lang.Byte.TYPE => classOf[java.lang.Byte] + case java.lang.Short.TYPE => classOf[java.lang.Short] + case java.lang.Character.TYPE => classOf[java.lang.Character] + case java.lang.Integer.TYPE => classOf[java.lang.Integer] + case java.lang.Long.TYPE => classOf[java.lang.Long] + case java.lang.Float.TYPE => classOf[java.lang.Float] + case java.lang.Double.TYPE => classOf[java.lang.Double] + case java.lang.Boolean.TYPE => classOf[java.lang.Boolean] + case java.lang.Void.TYPE => classOf[java.lang.Void] + case _ => clazz + } + } + + /** + * Convert a [[org.json4s.reflect.ScalaType]] to a [[com.fasterxml.jackson.databind.JavaType]]. + * This is used when defining the structure of a case class used in parsing JSON. + * + * @note order matters since Maps can be treated as collections of tuples + */ + def javaType( + typeFactory: TypeFactory, + scalaType: ScalaType + ): JavaType = { + // for backwards compatibility we box primitive types + val erasedClazzType = toScalaBoxedType(scalaType.erasure) + + if (scalaType.typeArgs.isEmpty || scalaType.erasure.isEnum) { + typeFactory.constructType(erasedClazzType) + } else if (scalaType.isCollection) { + if (scalaType.isMap) { + typeFactory.constructMapLikeType( + erasedClazzType, + javaType(typeFactory, scalaType.typeArgs.head), + javaType(typeFactory, scalaType.typeArgs(1)) + ) + } else if (scalaType.isArray) { + // need to special-case array creation + ArrayType.construct( + javaType(typeFactory, scalaType.typeArgs.head), + TypeBindings + .create( + classOf[ + java.util.ArrayList[_] + ], // we hardcode the type to `java.util.ArrayList` to properly support Array creation + javaType(typeFactory, scalaType.typeArgs.head) + ) + ) + } else { + typeFactory.constructCollectionLikeType( + erasedClazzType, + javaType(typeFactory, scalaType.typeArgs.head) + ) + } + } else { + typeFactory.constructParametricType( + erasedClazzType, + javaTypes(typeFactory, scalaType.typeArgs): _* + ) + } + } + + /** + * Convert a [[org.json4s.reflect.ScalaType]] to a [[com.fasterxml.jackson.databind.JavaType]] + * taking into account any parameterized types and type bindings. This is used when parsing + * JSON into a type. + * + * @note Match order matters since Maps can be treated as collections of tuples. + */ + def javaType( + typeFactory: TypeFactory, + scalaType: ScalaType, + typeParameters: Array[JavaType] + ): JavaType = { + + if (scalaType.typeArgs.isEmpty || scalaType.erasure.isEnum) { + typeParameters.head + } else if (scalaType.isCollection) { + if (scalaType.isMap) { + typeFactory.constructMapLikeType( + scalaType.erasure, + typeParameters.head, + typeParameters.last + ) + } else if (scalaType.isArray) { + // need to special-case array creation + // we hardcode the type to `java.util.ArrayList` to properly support Array creation + ArrayType.construct( + javaType(typeFactory, scalaType.typeArgs.head, typeParameters), + TypeBindings.create(classOf[java.util.ArrayList[_]], typeParameters.head) + ) + } else { + typeFactory.constructCollectionLikeType( + scalaType.erasure, + typeParameters.head + ) + } + } else { + typeFactory.constructParametricType(scalaType.erasure, typeParameters: _*) + } + } + + /* Private */ + + /* Primitive class to Scala wrapper (boxed) primitive class */ + private[this] def toScalaBoxedType(clazz: Class[_]): Class[_] = clazz match { + case java.lang.Byte.TYPE => classOf[Byte] + case java.lang.Short.TYPE => classOf[Short] + case java.lang.Character.TYPE => classOf[Char] + case java.lang.Integer.TYPE => classOf[Int] + case java.lang.Long.TYPE => classOf[Long] + case java.lang.Float.TYPE => classOf[Float] + case java.lang.Double.TYPE => classOf[Double] + case java.lang.Boolean.TYPE => classOf[Boolean] + case _ => clazz + } + + /* Create a Seq of JavaTypes from the given ScalaTypes */ + private[this] def javaTypes( + typeFactory: TypeFactory, + scalaTypes: Seq[ScalaType] + ): Seq[JavaType] = { + for (scalaType <- scalaTypes) yield { + javaType(typeFactory, scalaType) + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/BUILD b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/BUILD new file mode 100644 index 0000000000..2b2bdcf4a8 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/BUILD @@ -0,0 +1,22 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-caseclass-exceptions", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "util/util-validator/src/main/scala/com/twitter/util/validation", + ], + exports = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + ], +) diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassFieldMappingException.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassFieldMappingException.scala new file mode 100644 index 0000000000..6329853660 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassFieldMappingException.scala @@ -0,0 +1,103 @@ +package com.twitter.util.jackson.caseclass.exceptions + +import com.fasterxml.jackson.core.JsonProcessingException +import com.fasterxml.jackson.databind.JsonMappingException +import jakarta.validation.{ConstraintViolation, Payload} + +object CaseClassFieldMappingException { + + /** Marker trait for CaseClassFieldMappingException detail */ + sealed trait Detail + + /** No specified detail */ + case object Unspecified extends Detail + + /** A case class field specified as required is missing from JSON to deserialize */ + case object RequiredFieldMissing extends Detail + + object ValidationError { + sealed trait Location + case object Field extends Location + case object Method extends Location + } + + /** A violation was raised when performing a field or method validation */ + case class ValidationError( + violation: ConstraintViolation[_], + location: ValidationError.Location, + payload: Option[Payload]) + extends Detail + + /** A JsonProcessingException occurred during deserialization */ + case class JsonProcessingError(cause: JsonProcessingException) extends Detail + + /** Throwable detail which specifies a message */ + case class ThrowableError(message: String, cause: Throwable) extends Detail + + /** + * Represents the reason with message and more details for the field mapping exception. + */ + case class Reason( + message: String, + detail: Detail = Unspecified) + + object PropertyPath { + val Empty: PropertyPath = PropertyPath(Seq.empty) + private val FieldSeparator = "." + def leaf(name: String): PropertyPath = Empty.withParent(name) + } + + /** Represents a path to the case class field/property germane to the exception */ + case class PropertyPath(names: Seq[String]) { + private[jackson] def withParent(name: String): PropertyPath = copy(name +: names) + + def isEmpty: Boolean = names.isEmpty + def prettyString: String = names.mkString(PropertyPath.FieldSeparator) + } +} + +/** + * A subclass of [[JsonMappingException]] which bundles together a failed field location as a + * `CaseClassFieldMappingException.PropertyPath` with the the failure reason represented by a + * `CaseClassFieldMappingException.Reason` to carry the failure reason. + * + * @note this exception is a case class in order to have a useful equals() and hashcode() + * methods since this exception is generally carried in a collection inside of a + * [[CaseClassMappingException]]. + * @param path - a `CaseClassFieldMappingException.PropertyPath` instance to the case class field + * that caused the failure. + * @param reason - an instance of a `CaseClassFieldMappingException.Reason` which carries detail of the + * failure reason. + * + * @see [[com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.PropertyPath]] + * @see [[com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.Reason]] + * @see [[com.fasterxml.jackson.databind.JsonMappingException]] + * @see [[CaseClassMappingException]] + */ +case class CaseClassFieldMappingException( + path: CaseClassFieldMappingException.PropertyPath, + reason: CaseClassFieldMappingException.Reason) + extends JsonMappingException(null, reason.message) { + + /* Public */ + + /** + * Render a human readable message. If the error message pertains to + * a specific field it is prefixed with the field's name. + */ + override def getMessage: String = { + if (path == null || path.isEmpty) reason.message + else s"${path.prettyString}: ${reason.message}" + } + + /* Private */ + + /** fill in any missing PropertyPath information */ + private[jackson] def withPropertyPath( + path: CaseClassFieldMappingException.PropertyPath + ): CaseClassFieldMappingException = this.copy(path) + + private[jackson] def scoped(fieldName: String): CaseClassFieldMappingException = { + copy(path = path.withParent(fieldName)) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassMappingException.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassMappingException.scala new file mode 100644 index 0000000000..613fce0fd9 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/CaseClassMappingException.scala @@ -0,0 +1,77 @@ +package com.twitter.util.jackson.caseclass.exceptions + +import com.fasterxml.jackson.databind.JsonMappingException + +object CaseClassMappingException { + + /** + * Create a new [[CaseClassMappingException]] with no field exceptions. + * @return a new [[CaseClassMappingException]] + */ + def apply(): CaseClassMappingException = new CaseClassMappingException() + + /** + * Create a new [[CaseClassMappingException]] over the given field exceptions. + * @param fieldMappingExceptions exceptions to enclose in returned [[CaseClassMappingException]] + * @return a new [[CaseClassMappingException]] with the given field exceptions. + */ + def apply( + fieldMappingExceptions: Set[CaseClassFieldMappingException] + ): CaseClassMappingException = + new CaseClassMappingException(fieldMappingExceptions) +} + +/** + * A subclass of [[JsonMappingException]] used to signal fatal problems with mapping of JSON + * content to a Scala case class. + * + * Per-field details (of type [[CaseClassFieldMappingException]]) are carried to provide the + * ability to iterate over all exceptions causing the failure to construct the case class. + * + * This extends [[JsonMappingException]] such that this exception is properly handled + * when deserializing into nested case-classes. + * + * @see [[CaseClassFieldMappingException]] + * @see [[com.fasterxml.jackson.databind.JsonMappingException]] + */ +class CaseClassMappingException private ( + fieldMappingExceptions: Set[CaseClassFieldMappingException] = + Set.empty[CaseClassFieldMappingException]) + extends JsonMappingException(null, "") { + + /** + * The collection of [[CaseClassFieldMappingException]] instances which make up this + * [[CaseClassMappingException]]. This collection is intended to be purposely exhaustive in that + * is specifies all errors encountered in mapping JSON content to a Scala case class. + */ + val errors: Seq[CaseClassFieldMappingException] = + fieldMappingExceptions.toSeq.sortBy(_.getMessage) + + // helpers for formatting the exception message + private[this] val errorsSize = errors.size + private[this] val messagePreambleString = + if (errorsSize == 1) "An error was " else s"${errorsSize} errors " + private[this] val errorsString = + if (errorsSize == 1) "Error: " else "Errors: " + private[this] val errorsSeparatorString = "\n\t " + + /** + * Formats a human-readable message which includes the underlying [[CaseClassMappingException]] messages. + * + * ==Example== + * Multiple errors: + * {{{ + * 2 errors encountered during deserialization. + * Errors: com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException: data: must not be empty + * com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException: number: must be greater than or equal to 5 + * }}} + * Single error: + * {{{ + * An error was encountered during deserialization. + * Error: com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException: data: must not be empty + * }}} + */ + override def getMessage: String = + s"${messagePreambleString}encountered during deserialization.\n\t$errorsString" + + errors.mkString(errorsSeparatorString) +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/InjectableValuesException.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/InjectableValuesException.scala new file mode 100644 index 0000000000..03cbc9dcb7 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/InjectableValuesException.scala @@ -0,0 +1,22 @@ +package com.twitter.util.jackson.caseclass.exceptions + +import scala.util.control.NoStackTrace + +/** + * Represents an exception during deserialization due to the inability to properly inject a value via + * Jackson [[com.fasterxml.jackson.databind.InjectableValues]]. A common cause is an incorrectly + * configured mapper that is not instantiated with an appropriate instance of a + * [[com.fasterxml.jackson.databind.InjectableValues]] that can locate the requested value to inject + * during deserialization. + * + * @note this exception is not handled during case class deserialization and is thus expected to be + * handled by callers. + */ +class InjectableValuesException(message: String, cause: Throwable) + extends Exception(message, cause) + with NoStackTrace { + + def this(message: String) { + this(message, null) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/package.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/package.scala new file mode 100644 index 0000000000..841f0f5158 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions/package.scala @@ -0,0 +1,27 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.core.JsonProcessingException +import com.fasterxml.jackson.databind.JsonMappingException + +package object exceptions { + + implicit class RichJsonProcessingException[E <: JsonProcessingException](val self: E) + extends AnyVal { + + /** + * Converts a [[JsonProcessingException]] to a informative error message that can be returned + * to callers. The goal is to not leak internals of the underlying types. + * @return a useful [[String]] error message + */ + def errorMessage: String = { + self match { + case exc: JsonMappingException => + exc.getOriginalMessage // want to return the original message without and location reference + case _ if self.getCause == null => + self.getOriginalMessage // jackson threw the original error + case _ => + self.getCause.getMessage // custom deserialization code threw the exception (e.g., enum deserialization) + } + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/package.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/package.scala new file mode 100644 index 0000000000..8c85c36008 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/package.scala @@ -0,0 +1,148 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.annotation.JacksonAnnotation +import com.fasterxml.jackson.databind.deser.impl.ValueInjector +import com.fasterxml.jackson.databind.introspect.{ + AnnotatedMember, + AnnotatedParameter, + AnnotatedWithParams, + AnnotationMap, + TypeResolutionContext +} +import com.fasterxml.jackson.databind.{BeanProperty, DeserializationContext, JavaType, PropertyName} +import java.lang.annotation.Annotation +import org.json4s.reflect.ClassDescriptor + +package object caseclass { + + /** Resolves a subtype from a given optional class */ + private[twitter] def resolveSubType( + context: DeserializationContext, + baseType: JavaType, + subClazz: Option[Class[_]] + ): JavaType = { + subClazz match { + case Some(subTypeClazz) if isNotAssignableFrom(subTypeClazz, baseType.getRawClass) => + context.resolveSubType(baseType, subTypeClazz.getName) + case _ => + baseType + } + } + + /** Returns true if the class is not null and not assignable to "other" */ + private[twitter] def isNotAssignableFrom(clazz: Class[_], other: Class[_]): Boolean = { + clazz != null && !isAssignableFrom(other, clazz) + } + + /** boxes both sides for comparison -- no way to turn the from into the to type */ + private[twitter] def isAssignableFrom(fromType: Class[_], toType: Class[_]): Boolean = + Types.wrapperType(toType).isAssignableFrom(Types.wrapperType(fromType)) + + /** Create a new [[CaseClassBeanProperty]]. Returns the [[CaseClassBeanProperty]] and the JacksonInject.Value */ + private[twitter] def newBeanProperty( + context: DeserializationContext, + javaType: JavaType, + optionalJavaType: Option[JavaType], + annotatedParameter: AnnotatedParameter, + annotations: Iterable[Annotation], + name: String, + index: Int + ): CaseClassBeanProperty = { + // to support deserialization with JsonFormat and other Jackson annotations on fields, + // the Jackson annotations must be carried in the mutator of the ValueInjector. + val jacksonAnnotations = new AnnotationMap() + val contextAnnotationsMap = new AnnotationMap() + annotations.foreach { annotation => + if (annotation.annotationType().isAnnotationPresent(classOf[JacksonAnnotation])) { + jacksonAnnotations.add(annotation) + } else { + contextAnnotationsMap.add(annotation) + } + } + + val mutator: AnnotatedMember = optionalJavaType match { + case Some(optJavaType) => + newAnnotatedParameter( + context, + annotatedParameter, + AnnotationMap.merge(jacksonAnnotations, contextAnnotationsMap), + optJavaType, + index) + case _ => + newAnnotatedParameter( + context, + annotatedParameter, + AnnotationMap.merge(jacksonAnnotations, contextAnnotationsMap), + javaType, + index) + } + + val jacksonInjectValue = context.getAnnotationIntrospector.findInjectableValue(mutator) + val beanProperty: BeanProperty = new ValueInjector( + /* propName = */ new PropertyName(name), + /* type = */ optionalJavaType.getOrElse(javaType), + /* mutator = */ mutator, + /* valueId = */ if (jacksonInjectValue == null) null else jacksonInjectValue.getId + ) { + // ValueInjector no longer supports passing contextAnnotations as an + // argument as of jackson 2.9.x + // https://github.com/FasterXML/jackson-databind/commit/76381c528c9b75265e8b93bf6bb4532e4aa8e957#diff-dbcd29e987f27d95f964a674963a9066R24 + private[this] val contextAnnotations = contextAnnotationsMap + + override def getContextAnnotation[A <: Annotation](clazz: Class[A]): A = { + contextAnnotations.get[A](clazz) + } + } + + CaseClassBeanProperty(valueObj = Option(jacksonInjectValue), property = beanProperty) + } + + /** Create an [[AnnotatedMember]] for use in a [[CaseClassBeanProperty]] */ + private[jackson] def newAnnotatedParameter( + context: DeserializationContext, + base: AnnotatedParameter, + annotations: AnnotationMap, + javaType: JavaType, + index: Int + ): AnnotatedMember = + newAnnotatedParameter( + new TypeResolutionContext.Basic(context.getTypeFactory, javaType.getBindings), + if (base == null) null else base.getOwner, + annotations, + javaType, + index + ) + + /** Create an [[AnnotatedMember]] for use in a [[CaseClassBeanProperty]] */ + private[jackson] def newAnnotatedParameter( + typeResolutionContext: TypeResolutionContext, + owner: AnnotatedWithParams, + annotations: AnnotationMap, + javaType: JavaType, + index: Int + ): AnnotatedMember = { + new AnnotatedParameter( + /* owner = */ owner, + /* type = */ javaType, + /* typeResolutionContext = */ typeResolutionContext, + /* annotations = */ annotations, + /* index = */ index + ) + } + + private[jackson] def defaultMethod( + clazzDescriptor: ClassDescriptor, + constructorParamIdx: Int + ): Option[() => Object] = { + val defaultMethodArgNum = constructorParamIdx + 1 + for { + singletonDescriptor <- clazzDescriptor.companion + companionObjectClass = singletonDescriptor.erasure.erasure + companionObject = singletonDescriptor.instance + method <- companionObjectClass.getMethods.find( + _.getName == "$lessinit$greater$default$" + defaultMethodArgNum) + } yield { () => + method.invoke(companionObject) + } + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/package.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/package.scala new file mode 100644 index 0000000000..938625a6c8 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/package.scala @@ -0,0 +1,8 @@ +package com.twitter.util + +import com.fasterxml.jackson.databind.{ObjectMapper => JacksonObjectMapper} +import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} + +package object jackson { + type JacksonScalaObjectMapperType = JacksonObjectMapper with JacksonScalaObjectMapper +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/BUILD b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/BUILD new file mode 100644 index 0000000000..3d56d5356c --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jackson-serde", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-core:util-core-util", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions", + ], +) diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DefaultSerdeModule.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DefaultSerdeModule.scala new file mode 100644 index 0000000000..de1679e51b --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DefaultSerdeModule.scala @@ -0,0 +1,12 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.databind.module.SimpleModule +import com.twitter.{util => ctu} + +private[jackson] object DefaultSerdeModule extends SimpleModule { + addSerializer(WrappedValueSerializer) + addSerializer(DurationStringSerializer) + addSerializer(TimeStringSerializer()) + addDeserializer(classOf[ctu.Duration], DurationStringDeserializer) + addDeserializer(classOf[ctu.Time], TimeStringDeserializer()) +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringDeserializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringDeserializer.scala new file mode 100644 index 0000000000..7b63d7b689 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringDeserializer.scala @@ -0,0 +1,43 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.core.{JsonParser, JsonToken} +import com.fasterxml.jackson.databind.{ + DeserializationContext, + JsonDeserializer, + JsonMappingException +} +import com.twitter.util.Duration +import scala.util.control.NonFatal + +/** + * A Jackson [[JsonDeserializer]] for [[com.twitter.util.Duration]]. + */ +object DurationStringDeserializer extends JsonDeserializer[Duration] { + + override final def deserialize(jp: JsonParser, ctxt: DeserializationContext): Duration = + jp.currentToken match { + case JsonToken.NOT_AVAILABLE | JsonToken.START_OBJECT => + handleToken(jp.nextToken, jp, ctxt) + case _ => + handleToken(jp.currentToken, jp, ctxt) + } + + private[this] def handleToken( + token: JsonToken, + jp: JsonParser, + ctxt: DeserializationContext + ): Duration = { + val value = jp.getText.trim + if (value.isEmpty) { + throw JsonMappingException.from(ctxt, "field cannot be empty") + } else { + try { + Duration.parse(value) + } catch { + case NonFatal(e) => + throw JsonMappingException.from(ctxt, e.getMessage, e) + } + } + } + +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringSerializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringSerializer.scala new file mode 100644 index 0000000000..a6d9517a67 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/DurationStringSerializer.scala @@ -0,0 +1,20 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.core.JsonGenerator +import com.fasterxml.jackson.databind.SerializerProvider +import com.fasterxml.jackson.databind.ser.std.StdSerializer +import com.twitter.util.Duration + +/** + * A Jackson [[StdSerializer]] for [[com.twitter.util.Duration]]. + */ +object DurationStringSerializer extends StdSerializer[Duration](classOf[Duration]) { + + override def serialize( + value: Duration, + jgen: JsonGenerator, + provider: SerializerProvider + ): Unit = { + jgen.writeString(value.toString) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/LongKeyDeserializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/LongKeyDeserializer.scala new file mode 100644 index 0000000000..42ac80fd39 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/LongKeyDeserializer.scala @@ -0,0 +1,52 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.databind._ +import com.fasterxml.jackson.databind.deser.KeyDeserializers +import com.fasterxml.jackson.module.scala.JacksonModule +import com.twitter.util.WrappedValue +import org.json4s.reflect.{ClassDescriptor, Reflector, classDescribable} + +private[jackson] class LongKeyDeserializer(clazz: Class[_]) extends KeyDeserializer { + private val constructor = clazz.getConstructor(classOf[Long]) + + override def deserializeKey(key: String, ctxt: DeserializationContext): Object = { + val long = key.toLong.asInstanceOf[Object] + constructor.newInstance(long).asInstanceOf[Object] + } +} + +private[jackson] object LongKeyDeserializers extends JacksonModule { + override def getModuleName = "LongKeyDeserializers" + + private val keyDeserializers = new KeyDeserializers { + override def findKeyDeserializer( + `type`: JavaType, + config: DeserializationConfig, + beanDesc: BeanDescription + ): KeyDeserializer = { + val clazz = beanDesc.getBeanClass + if (isJsonWrappedLong(clazz)) + new LongKeyDeserializer(clazz) + else + null + } + } + + private def isJsonWrappedLong(clazz: Class[_]): Boolean = { + classOf[WrappedValue[_]].isAssignableFrom(clazz) && + isWrappedLong(clazz) + } + + private def isWrappedLong(clazz: Class[_]): Boolean = { + val reflector = + Reflector + .describe(classDescribable(clazz)) + .asInstanceOf[ClassDescriptor] + .constructors + .head + val erased = reflector.params.head.argType.erasure + erased == classOf[Long] || erased.getName == "long" + } + + this += { _.addKeyDeserializers(keyDeserializers) } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringDeserializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringDeserializer.scala new file mode 100644 index 0000000000..cd7b54f393 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringDeserializer.scala @@ -0,0 +1,99 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.core.{JsonParser, JsonToken} +import com.fasterxml.jackson.databind.deser.ContextualDeserializer +import com.fasterxml.jackson.databind.deser.std.StdScalarDeserializer +import com.fasterxml.jackson.databind.{ + BeanProperty, + DeserializationContext, + JsonDeserializer, + JsonMappingException +} +import com.twitter.util.{Time, TimeFormat} +import java.util.{Locale, TimeZone} +import scala.util.control.NonFatal + +object TimeStringDeserializer { + private[this] val DefaultTimeFormat: String = "yyyy-MM-dd'T'HH:mm:ss.SSSZ" + + def apply(): TimeStringDeserializer = this(DefaultTimeFormat) + + def apply(pattern: String): TimeStringDeserializer = + this(pattern, None, TimeZone.getTimeZone("UTC")) + + def apply(pattern: String, locale: Option[Locale], timezone: TimeZone): TimeStringDeserializer = + new TimeStringDeserializer(new TimeFormat(pattern, locale, timezone)) +} + +/** + * A Jackson [[JsonDeserializer]] for [[com.twitter.util.Time]]. + * @param timeFormat the configured [[com.twitter.util.TimeFormat]] for this deserializer. + */ +class TimeStringDeserializer(private[this] val timeFormat: TimeFormat) + extends StdScalarDeserializer[Time](classOf[Time]) + with ContextualDeserializer { + + override def deserialize(jp: JsonParser, ctxt: DeserializationContext): Time = + jp.currentToken match { + case JsonToken.NOT_AVAILABLE | JsonToken.START_OBJECT => + handleToken(jp.nextToken, jp, ctxt) + case _ => + handleToken(jp.currentToken, jp, ctxt) + } + + private[this] def handleToken( + token: JsonToken, + jp: JsonParser, + ctxt: DeserializationContext + ): Time = { + try { + val value = jp.getText.trim + if (value.isEmpty) { + throw JsonMappingException.from(ctxt, "field cannot be empty") + } else { + timeFormat.parse(value) + } + } catch { + case NonFatal(e) => + throw JsonMappingException.from( + ctxt, + "error parsing '" + jp.getText + s"' into a ${Time.getClass.getName}", + e) + } + } + + /** + * This method allows extracting the [[com.fasterxml.jackson.annotation.JsonFormat JsonFormat]] + * annotation to create a [[com.twitter.util.TimeFormat TimeFormat]] based on the specifications + * provided in the annotation. The implementation follows the Jackson's java8 & joda-time versions + * + * @param ctxt Deserialization context to access configuration, additional deserializers + * that may be needed by this deserializer + * @param property Method, field or constructor parameter that represents the property (and is + * used to assign deserialized value). Should be available; but there may be + * cases where caller can not provide it and null is passed instead + * (in which case impls usually pass 'this' deserializer as is) + * + * @return Deserializer to use for deserializing values of specified property; may be this + * instance or a new instance. + * + * @see https://github.com/FasterXML/jackson-modules-java8/blob/master/datetime/src/main/java/com/fasterxml/jackson/datatype/jsr310/deser/JSR310DateTimeDeserializerBase.java#L29 + */ + override def createContextual( + ctxt: DeserializationContext, + property: BeanProperty + ): JsonDeserializer[_] = { + val deserializerOption: Option[TimeStringDeserializer] = for { + jsonFormat <- Option(findFormatOverrides(ctxt, property, handledType())) + deserializer <- Option( + TimeStringDeserializer( + jsonFormat.getPattern, + Option(jsonFormat.getLocale), + Option(jsonFormat.getTimeZone).getOrElse(TimeZone.getTimeZone("UTC")) + )) if jsonFormat.hasPattern + } yield { + deserializer + } + deserializerOption.getOrElse(this) + } +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringSerializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringSerializer.scala new file mode 100644 index 0000000000..e621dec917 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/TimeStringSerializer.scala @@ -0,0 +1,73 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.core.JsonGenerator +import com.fasterxml.jackson.databind.ser.ContextualSerializer +import com.fasterxml.jackson.databind.ser.std.StdScalarSerializer +import com.fasterxml.jackson.databind.{BeanProperty, JsonSerializer, SerializerProvider} +import com.twitter.util.{Time, TimeFormat} +import java.util.{Locale, TimeZone} + +object TimeStringSerializer { + // TODO: this should be symmetric to the deserializer but we introduced this originally + // without contextual information and thus all strings were serialized with the default + // format of the object mapper ("yyyy-MM-dd HH:mm:ss Z"). This will require a migration. + // private[this] val DefaultTimeFormat = "yyyy-MM-dd'T'HH:mm:ss.SSSZ" + private[this] val DefaultTimeFormat = "yyyy-MM-dd HH:mm:ss Z" + + def apply(): TimeStringSerializer = this(DefaultTimeFormat) + + def apply(pattern: String): TimeStringSerializer = + this(pattern, None, TimeZone.getTimeZone("UTC")) + + def apply(pattern: String, locale: Option[Locale], timezone: TimeZone): TimeStringSerializer = + new TimeStringSerializer(new TimeFormat(pattern, locale, timezone)) +} + +/** + * A Jackson [[JsonSerializer]] for [[com.twitter.util.Time]]. + * @param timeFormat the configured [[com.twitter.util.TimeFormat]] for this serializer. + */ +class TimeStringSerializer(private[this] val timeFormat: TimeFormat) + extends StdScalarSerializer[Time](classOf[Time]) + with ContextualSerializer { + + override def serialize(value: Time, jgen: JsonGenerator, provider: SerializerProvider): Unit = { + jgen.writeString(timeFormat.format(value)) + } + + /** + * This method allows extracting the [[com.fasterxml.jackson.annotation.JsonFormat JsonFormat]] + * annotation and create a [[com.twitter.util.TimeFormat TimeFormat]] based on the specifications + * provided in the annotation. The implementation follows the Jackson's java8 & joda-time versions + * + * @param provider Serialization provider to access configuration, additional deserializers + * that may be needed by this deserializer + * @param property Method, field or constructor parameter that represents the property (and is + * used to assign serialized value). Should be available; but there may be + * cases where caller can not provide it and null is passed instead + * (in which case impls usually pass 'this' serializer as is) + * + * @return Serializer to use for serializing values of specified property; may be this + * instance or a new instance. + * + * @see https://github.com/FasterXML/jackson-modules-java8/blob/master/datetime/src/main/java/com/fasterxml/jackson/datatype/jsr310/ser/JSR310FormattedSerializerBase.java#L114 + */ + override def createContextual( + provider: SerializerProvider, + property: BeanProperty + ): JsonSerializer[_] = { + val serializerOption: Option[TimeStringSerializer] = for { + jsonFormat <- Option(findFormatOverrides(provider, property, handledType())) + serializer <- Option( + TimeStringSerializer( + jsonFormat.getPattern, + Option(jsonFormat.getLocale), + Option(jsonFormat.getTimeZone).getOrElse(TimeZone.getTimeZone("UTC")) + )) if jsonFormat.hasPattern + } yield { + serializer + } + serializerOption.getOrElse(this) + } + +} diff --git a/util-jackson/src/main/scala/com/twitter/util/jackson/serde/WrappedValueSerializer.scala b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/WrappedValueSerializer.scala new file mode 100644 index 0000000000..957e7e2fd6 --- /dev/null +++ b/util-jackson/src/main/scala/com/twitter/util/jackson/serde/WrappedValueSerializer.scala @@ -0,0 +1,18 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.core.JsonGenerator +import com.fasterxml.jackson.databind.SerializerProvider +import com.fasterxml.jackson.databind.ser.std.StdSerializer +import com.twitter.util.WrappedValue + +private[jackson] object WrappedValueSerializer + extends StdSerializer[WrappedValue[_]](classOf[WrappedValue[_]]) { + + override def serialize( + wrappedValue: WrappedValue[_], + jgen: JsonGenerator, + provider: SerializerProvider + ): Unit = { + jgen.writeObject(wrappedValue.onlyValue) + } +} diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/BUILD.bazel b/util-jackson/src/test/java/com/twitter/util/jackson/BUILD.bazel new file mode 100644 index 0000000000..e834d5832e --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/BUILD.bazel @@ -0,0 +1,10 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], +) diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/CarMake.java b/util-jackson/src/test/java/com/twitter/util/jackson/CarMake.java new file mode 100644 index 0000000000..6dbc473047 --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/CarMake.java @@ -0,0 +1,8 @@ +package com.twitter.util.jackson; + +/** + * Test Java Enum + */ +public enum CarMake { + Ford, Honda; +} diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/CarMakeEnum.java b/util-jackson/src/test/java/com/twitter/util/jackson/CarMakeEnum.java new file mode 100644 index 0000000000..800d464092 --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/CarMakeEnum.java @@ -0,0 +1,9 @@ +package com.twitter.util.jackson; + +/** + * Test Java Enum + */ +public enum CarMakeEnum { + ford, + vw; +} diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/TestInjectableValue.java b/util-jackson/src/test/java/com/twitter/util/jackson/TestInjectableValue.java new file mode 100644 index 0000000000..adbfb60f47 --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/TestInjectableValue.java @@ -0,0 +1,19 @@ +package com.twitter.util.jackson; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * FOR TESTING ONLY + */ +@Target(ElementType.PARAMETER) +@Retention(RetentionPolicy.RUNTIME) +public @interface TestInjectableValue { + /** + * FOR TESTING ONLY + * @return the value + */ + String value() default ""; +} diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/tests/BUILD.bazel b/util-jackson/src/test/java/com/twitter/util/jackson/tests/BUILD.bazel new file mode 100644 index 0000000000..d662bae80b --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/tests/BUILD.bazel @@ -0,0 +1,12 @@ +junit_tests( + sources = ["**/*.java"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/dataformat:jackson-dataformat-yaml", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-jackson/src/main/scala/com/twitter/util/jackson", + ], +) diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/tests/JsonDiffJavaTest.java b/util-jackson/src/test/java/com/twitter/util/jackson/tests/JsonDiffJavaTest.java new file mode 100644 index 0000000000..492c36c88b --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/tests/JsonDiffJavaTest.java @@ -0,0 +1,86 @@ +package com.twitter.util.jackson.tests; + +import java.io.ByteArrayOutputStream; +import java.io.PrintStream; + +import scala.Option; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.node.ObjectNode; + +import org.junit.Assert; +import org.junit.Test; + +import com.twitter.util.jackson.JsonDiff; + +public class JsonDiffJavaTest extends Assert { + + @Test + public void testDiff() { + String a = "{\"a\": 1, \"b\": 2}"; + String b = "{\"b\": 2, \"a\": 1}"; + + Option result = JsonDiff.diff(a, b); + assertFalse(result.isDefined()); + } + + @Test + public void testDiffWithNormalizer() { + String a = "{\"a\": 1, \"b\": 2, \"c\": 5}"; + String b = "{\"b\": 2, \"a\": 1, \"c\": 3}"; + + // normalizerFn will "normalize" the value for "c" into the expected value + Option result = JsonDiff.diff(a, b, (JsonNode jsonNode) -> ((ObjectNode)jsonNode).put("c", 5) ); + assertFalse(result.isDefined()); + } + + @Test + public void testAssertDiff() { + String a = "{\"a\": 1, \"b\": 2}"; + String b = "{\"b\": 2, \"a\": 1}"; + + JsonDiff.assertDiff(a, b); + } + + @Test + public void testAssertDiffWithNormalizer() { + String a = "{\"a\": 1, \"b\": 2, \"c\": 5}"; + String b = "{\"b\": 2, \"a\": 1, \"c\": 3}"; + + // normalizerFn will "normalize" the value for "c" into the expected value + JsonDiff.assertDiff(a, b, (JsonNode jsonNode) -> ((ObjectNode)jsonNode).put("c", 5) ); + } + + @Test + public void testAssertDiffWithPrintStream() throws Exception { + String a = "{\"a\": 1, \"b\": 2}"; + String b = "{\"b\": 2, \"a\": 1}"; + + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + PrintStream ps = new PrintStream(baos, true, "utf-8"); + + try { + JsonDiff.assertDiff(a, b, ps); + } finally { + ps.flush(); + ps.close(); + } + } + + @Test + public void testAssertDiffWithNormalizerAndPrintStream() throws Exception { + String a = "{\"a\": 1, \"b\": 2, \"c\": 5}"; + String b = "{\"b\": 2, \"a\": 1, \"c\": 3}"; + + ByteArrayOutputStream baos = new ByteArrayOutputStream(); + PrintStream ps = new PrintStream(baos, true, "utf-8"); + + try { + // normalizerFn will "normalize" the value for "c" into the expected value + JsonDiff.assertDiff(a, b, (JsonNode jsonNode) -> ((ObjectNode)jsonNode).put("c", 5) ); + } finally { + ps.flush(); + ps.close(); + } + } +} diff --git a/util-jackson/src/test/java/com/twitter/util/jackson/tests/ScalaObjectMapperJavaTest.java b/util-jackson/src/test/java/com/twitter/util/jackson/tests/ScalaObjectMapperJavaTest.java new file mode 100644 index 0000000000..4a09680646 --- /dev/null +++ b/util-jackson/src/test/java/com/twitter/util/jackson/tests/ScalaObjectMapperJavaTest.java @@ -0,0 +1,47 @@ +package com.twitter.util.jackson.tests; + +import com.fasterxml.jackson.core.JsonFactory; +import com.fasterxml.jackson.dataformat.yaml.YAMLFactory; + +import org.junit.Assert; +import org.junit.Test; + +import com.twitter.util.jackson.ScalaObjectMapper; + +public class ScalaObjectMapperJavaTest extends Assert { + + /* Class under test */ + private final ScalaObjectMapper mapper = ScalaObjectMapper.apply(); + + @Test + public void testConstructors() { + // new mapper with defaults + Assert.assertNotNull(ScalaObjectMapper.apply()); + + // augments the underlying mapper with the defaults + Assert.assertNotNull(ScalaObjectMapper.apply(mapper.underlying())); + + // copies the underlying Jackson object mapper + Assert.assertNotNull(ScalaObjectMapper.objectMapper(mapper.underlying())); + try { + ScalaObjectMapper.yamlObjectMapper(mapper.underlying()); + fail(); + } catch (IllegalArgumentException e) { + // underlying mapper needs to have a JsonFactory of type YAMLFactory + }; + Assert.assertNotNull(ScalaObjectMapper.yamlObjectMapper(ScalaObjectMapper.builder().yamlObjectMapper().underlying())); + Assert.assertNotNull(ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying())); + Assert.assertNotNull(ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying())); +} + + @Test + public void testBuilderConstructors() { + Assert.assertNotNull(ScalaObjectMapper.builder().objectMapper()); + Assert.assertNotNull(ScalaObjectMapper.builder().objectMapper(new JsonFactory())); + Assert.assertNotNull(ScalaObjectMapper.builder().objectMapper(new YAMLFactory())); + + Assert.assertNotNull(ScalaObjectMapper.builder().yamlObjectMapper()); + Assert.assertNotNull(ScalaObjectMapper.builder().camelCaseObjectMapper()); + Assert.assertNotNull(ScalaObjectMapper.builder().snakeCaseObjectMapper()); + } +} diff --git a/util-jackson/src/test/resources/BUILD.bazel b/util-jackson/src/test/resources/BUILD.bazel new file mode 100644 index 0000000000..301b666c21 --- /dev/null +++ b/util-jackson/src/test/resources/BUILD.bazel @@ -0,0 +1,4 @@ +resources( + sources = ["**/*"], + tags = ["bazel-compatible"], +) diff --git a/util-jackson/src/test/resources/test.json b/util-jackson/src/test/resources/test.json new file mode 100644 index 0000000000..98e5dbd51e --- /dev/null +++ b/util-jackson/src/test/resources/test.json @@ -0,0 +1,3 @@ +{ + "id": "55555555" +} diff --git a/util-jackson/src/test/resources/test.yml b/util-jackson/src/test/resources/test.yml new file mode 100644 index 0000000000..ac1f85b4b0 --- /dev/null +++ b/util-jackson/src/test/resources/test.yml @@ -0,0 +1,2 @@ +--- +id: 55555555 diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/BUILD.bazel b/util-jackson/src/test/scala/com/twitter/util/jackson/BUILD.bazel new file mode 100644 index 0000000000..b8221c4a97 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/BUILD.bazel @@ -0,0 +1,45 @@ +junit_tests( + sources = [ + "*.java", + "*.scala", + "caseclass/*.scala", + "serde/*.scala", + ], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-annotations", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/dataformat:jackson-dataformat-yaml", + "3rdparty/jvm/com/fasterxml/jackson/datatype:jackson-datatype-jsr310", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/json4s:json4s-core", + "3rdparty/jvm/org/scala-lang/modules:scala-collection-compat", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "3rdparty/jvm/org/slf4j:slf4j-api", + scoped( + "3rdparty/jvm/org/slf4j:slf4j-simple", + scope = "runtime", + ), + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-jackson/src/main/scala/com/fasterxml/jackson/databind", + "util/util-jackson/src/main/scala/com/twitter/util/jackson", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/caseclass/exceptions", + "util/util-jackson/src/main/scala/com/twitter/util/jackson/serde", + "util/util-jackson/src/test/java/com/twitter/util/jackson", + "util/util-jackson/src/test/resources", + "util/util-mock/src/main/scala/com/twitter/util/mock", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation", + ], +) diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/FileResources.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/FileResources.scala new file mode 100644 index 0000000000..d68f8fc6b5 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/FileResources.scala @@ -0,0 +1,29 @@ +package com.twitter.util.jackson + +import com.twitter.io.TempFolder +import java.io.{FileOutputStream, OutputStream, File => JFile} +import java.nio.charset.StandardCharsets +import java.nio.file.{FileSystems, Files} + +trait FileResources extends TempFolder { + + protected def writeStringToFile( + directory: String, + name: String, + ext: String, + data: String + ): JFile = { + val file = Files.createTempFile(FileSystems.getDefault.getPath(directory), name, ext).toFile + try { + val out: OutputStream = new FileOutputStream(file, false) + out.write(data.getBytes(StandardCharsets.UTF_8)) + file + } finally { + file.deleteOnExit() + } + } + + protected def addSlash(directory: String): String = { + if (directory.endsWith("/")) directory else s"$directory/" + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/JSONTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/JSONTest.scala new file mode 100644 index 0000000000..c7af26b123 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/JSONTest.scala @@ -0,0 +1,120 @@ +package com.twitter.util.jackson + +import com.twitter.io.Buf +import java.io.{ByteArrayInputStream, File => JFile} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class JSONTest extends AnyFunSuite with Matchers with FileResources { + + test("JSON#parse 1") { + JSON.parse[Map[String, Int]]("""{"a": 1, "b": 2}""") match { + case Some(map) => + map should be(Map("a" -> 1, "b" -> 2)) + case _ => fail() + } + } + + test("JSON#parse 2") { + JSON.parse[Seq[String]]("""["a", "b", "c"]""") match { + case Some(Seq("a", "b", "c")) => // pass + case _ => fail() + } + } + + test("JSON#parse 3") { + JSON.parse[FooClass]("""{"id": "abcd1234"}""") match { + case Some(foo) => + foo.id should equal("abcd1234") + case _ => fail() + } + } + + test("JSON#parse 4") { + JSON.parse[CaseClassWithNotEmptyValidation]("""{"name": "", "make": "vw"}""") match { + case Some(clazz) => + // performs no validations + clazz.name.isEmpty should be(true) + clazz.make should equal(CarMakeEnum.vw) + case _ => fail() + } + } + + test("JSON#parse 5") { + val buf = Buf.Utf8("""{"id": "abcd1234"}""") + JSON.parse[FooClass](buf) match { + case Some(foo) => + foo.id should equal("abcd1234") + case _ => fail() + } + } + + test("JSON#parse 6") { + val inputStream = new ByteArrayInputStream("""{"id": "abcd1234"}""".getBytes("UTF-8")) + try { + JSON.parse[FooClass](inputStream) match { + case Some(foo) => + foo.id should equal("abcd1234") + case _ => fail() + } + } finally { + inputStream.close() + } + } + + test("JSON#parse 7") { + withTempFolder { + val file: JFile = + writeStringToFile(folderName, "test-file", ".json", """{"id": "999999999"}""") + JSON.parse[FooClass](file) match { + case Some(foo) => + foo.id should equal("999999999") + case _ => fail() + } + } + } + + test("JSON#write 1") { + JSON.write(Map("a" -> 1, "b" -> 2)) should equal("""{"a":1,"b":2}""") + } + + test("JSON#write 2") { + JSON.write(Seq("a", "b", "c")) should equal("""["a","b","c"]""") + } + + test("JSON#write 3") { + JSON.write(FooClass("abcd1234")) should equal("""{"id":"abcd1234"}""") + } + + test("JSON#prettyPrint 1") { + JSON.prettyPrint(Map("a" -> 1, "b" -> 2)) should equal("""{ + | "a" : 1, + | "b" : 2 + |}""".stripMargin) + } + + test("JSON#prettyPrint 2") { + JSON.prettyPrint(Seq("a", "b", "c")) should equal("""[ + | "a", + | "b", + | "c" + |]""".stripMargin) + } + + test("JSON#prettyPrint 3") { + JSON.prettyPrint(FooClass("abcd1234")) should equal("""{ + | "id" : "abcd1234" + |}""".stripMargin) + } + + test("JSON.Resource#parse resource") { + JSON.Resource.parse[FooClass]("/test.json") match { + case Some(foo) => + foo.id should equal("55555555") + case _ => fail() + } + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/JsonDiffTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/JsonDiffTest.scala new file mode 100644 index 0000000000..61e75f4a84 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/JsonDiffTest.scala @@ -0,0 +1,303 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.databind.JsonNode +import com.fasterxml.jackson.databind.node.ObjectNode +import java.io.{ByteArrayOutputStream, PrintStream} +import org.junit.runner.RunWith +import org.scalacheck.Gen +import org.scalatest.{BeforeAndAfterAll, BeforeAndAfterEach} +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import scala.util.control.NonFatal + +@RunWith(classOf[JUnitRunner]) +class JsonDiffTest + extends AnyFunSuite + with Matchers + with BeforeAndAfterEach + with BeforeAndAfterAll + with JsonGenerator + with ScalaCheckDrivenPropertyChecks { + + private[this] val mapper = ScalaObjectMapper.builder.objectMapper + + private[this] val baos = new ByteArrayOutputStream() + private[this] val ps = new PrintStream(baos, true, "utf-8") + + override protected def beforeEach(): Unit = { + ps.flush() + } + + override protected def afterAll(): Unit = { + try { + ps.close() // closes underlying outputstream + } catch { + case NonFatal(_) => // do nothing + } + super.afterAll() + } + + test("JsonDiff#diff success") { + val a = + """ + { + "a": 1, + "b": 2 + } + """ + + val b = + """ + { + "b": 2, + "a": 1 + } + """ + + val result = JsonDiff.diff(a, b) + result.isDefined should be(false) + } + + test("JsonDiff#diff with normalizer success") { + val expected = + """ + { + "a": 1, + "b": 2, + "c": 5 + } + """ + + // normalizerFn will "normalize" the value for "c" into the expected value + val actual = + """ + { + "b": 2, + "a": 1, + "c": 3 + } + """ + + val result = JsonDiff.diff( + expected, + actual, + normalizeFn = { jsonNode: JsonNode => + jsonNode.asInstanceOf[ObjectNode].put("c", 5) + } + ) + result.isDefined should be(false) + } + + test("JsonDiff#diff failure") { + val expected = + """ + { + "a": 1, + "b": 2 + } + """ + + val actual = + """ + { + "b": 22, + "a": 1 + } + """ + + val result = JsonDiff.diff(expected, actual) + result.isDefined should be(true) + } + + test("JsonDiff#diff with normalizer failure") { + val expected = + """ + { + "a": 1, + "b": 2, + "c": 5 + } + """ + + // normalizerFn only touches the value for "c" and "b" still doesn't match + val actual = + """ + { + "b": 22, + "a": 1, + "c": 3 + } + """ + + val result = JsonDiff.diff( + expected, + actual, + normalizeFn = { jsonNode: JsonNode => + jsonNode.asInstanceOf[ObjectNode].put("c", 5) + } + ) + result.isDefined should be(true) + } + + test("JsonDiff#assertDiff pass") { + val expected = + """ + { + "a": 1, + "b": 2 + } + """ + + val actual = + """ + { + "a": 1, + "b": 2 + } + """ + + JsonDiff.assertDiff(expected, actual, p = ps) + } + + test("JsonDiff#assertDiff with normalizer pass") { + val expected = + """ + { + "a": 1, + "b": 2, + "c": 3 + } + """ + + // normalizerFn will "normalize" the value for "c" into the expected value + val actual = + """ + { + "c": 5, + "a": 1, + "b": 2 + } + """ + + JsonDiff.assertDiff( + expected, + actual, + normalizeFn = { jsonNode: JsonNode => + jsonNode.asInstanceOf[ObjectNode].put("c", 3) + }, + p = ps) + } + + test("JsonDiff#assertDiff with normalizer fail") { + val expected = + """ + { + "a": 1, + "b": 2, + "c": 3 + } + """ + + // normalizerFn only touches the value for "c" and "a" still doesn't match + val actual = + """ + { + "c": 5, + "a": 11, + "b": 2 + } + """ + + intercept[AssertionError] { + JsonDiff.assertDiff( + expected, + actual, + normalizeFn = { jsonNode: JsonNode => + jsonNode.asInstanceOf[ObjectNode].put("c", 3) + }, + p = ps) + } + } + + test("JsonDiff#diff arbitrary pass") { + val expected = for { + depth <- Gen.chooseNum[Int](1, 10) + on <- objectNode(depth) + } yield on + + forAll(expected) { node => + val json = mapper.writeValueAsString(node) + JsonDiff.diff(expected = json, actual = json).isEmpty should be(true) + } + } + + test("JsonDiff#diff arbitrary fail") { + val values = for { + depth <- Gen.chooseNum[Int](1, 5) + expected <- objectNode(depth) + actual <- objectNode(depth) + if expected != actual + } yield (expected, actual) + + forAll(values) { + case (expected, actual) => + val result = JsonDiff.diff(expected = expected, actual = actual) + result.isDefined should be(true) + } + } + + // (the arbitrary tests above do not generate tests for array depth 1) + + test("JsonDiff#diff array depth 1 fail") { + val expected = """["a","b"]""" + val actual = """["b","a"]""" + JsonDiff.diff(expected, actual).isEmpty should be(false) + } + + test("JsonDiff#diff array depth 1 pass") { + val expected = """["a","b"]""" + val actual = """["a","b"]""" + JsonDiff.diff(expected, actual).isEmpty should be(true) + } + + // (the arbitrary tests above do not generate tests with null values or escapes) + + test("JsonDiff#diff with null values - fail") { + val expected = """{"a":null}""" + val actual = "{}" + JsonDiff.diff(expected, actual).isEmpty should be(false) + } + + test("JsonDiff#diff with null values - pass") { + val expected = """{"a":null}""" + val actual = """{"a": null}""" + JsonDiff.diff(expected, actual).isEmpty should be(true) + } + + test("JsonDiff#diff with unicode escape - pass") { + val expected = """{"t1": "24\u00B0F"}""" + val actual = """{"t1": "24°F"}""" + JsonDiff.diff(expected, actual).isEmpty should be(true) + } + + test("JsonDiff#toSortedString is sorted") { + val before = mapper.parse[JsonNode]("""{"a":1,"c":3,"b":2}""") + val expected = """{"a":1,"b":2,"c":3}""" + JsonDiff.toSortedString(before) should equal(expected) + } + + test("JsonDiff#toSortedString doesn't sort arrays") { + val before = mapper.parse[JsonNode]("""["b","a"]""") + val expected = """["b","a"]""" + JsonDiff.toSortedString(before) should equal(expected) + } + + test("JsonDiff#toSortedString includes null values") { + val before = mapper.parse[JsonNode]("""{"a":null}""") + val expected = """{"a":null}""" + JsonDiff.toSortedString(before) should equal(expected) + } + +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/JsonGenerator.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/JsonGenerator.scala new file mode 100644 index 0000000000..9ab0d334b7 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/JsonGenerator.scala @@ -0,0 +1,72 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.databind.JsonNode +import com.fasterxml.jackson.databind.node.{ + ArrayNode, + IntNode, + JsonNodeFactory, + ObjectNode, + TextNode +} +import org.scalacheck.Gen._ +import org.scalacheck._ +import scala.jdk.CollectionConverters._ + +object JsonGenerator extends JsonGenerator + +trait JsonGenerator { + + private[this] val jsonNodeFactory: JsonNodeFactory = new JsonNodeFactory(false) + + /** Generate either a ArrayNode or a ObjectNode */ + def jsonNode(depth: Int): Gen[JsonNode] = oneOf(arrayNode(depth), objectNode(depth)) + + /** Generate a ArrayNode */ + def arrayNode(depth: Int): Gen[ArrayNode] = for { + n <- choose(1, 4) + vals <- values(n, depth) + } yield new ArrayNode(jsonNodeFactory).addAll(vals.asJava) + + /** Generate a ObjectNode */ + def objectNode(depth: Int): Gen[ObjectNode] = for { + n <- choose(1, 4) + ks <- keys(n) + vals <- values(n, depth) + } yield { + val kids: Map[String, JsonNode] = Map(ks.zip(vals): _*) + new ObjectNode(jsonNodeFactory, kids.asJava) + } + + /** Generate a list of keys to be used in the map of a ObjectNode */ + private[this] def keys(n: Int): Gen[List[String]] = { + val key: Gen[String] = for { + str <- alphaStr.filter(_.nonEmpty).map(_.toLowerCase) + } yield { + if (str.length > 4) str.substring(0, 3) + else str + } + listOfN(n, key) + } + + /** + * Generate a list of values to be used in the map of a ObjectNode or in the list of an ArrayNode. + */ + private[this] def values(n: Int, depth: Int): Gen[List[JsonNode]] = listOfN(n, value(depth)) + + /** + * Generate a value to be used in the map of a ObjectNode or in the list of an ArrayNode. + */ + private[this] def value(depth: Int): Gen[JsonNode] = + if (depth == 0) terminalNode + else oneOf(jsonNode(depth - 1), terminalNode) + + /** Generate a terminal node */ + private[this] def terminalNode: Gen[JsonNode] = + oneOf( + new IntNode(4), + new IntNode(2), + new TextNode("b"), + new TextNode("i"), + new TextNode("r"), + new TextNode("d")) +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/NoValidatorImplConstraint.java b/util-jackson/src/test/scala/com/twitter/util/jackson/NoValidatorImplConstraint.java new file mode 100644 index 0000000000..61ae4d3398 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/NoValidatorImplConstraint.java @@ -0,0 +1,38 @@ +package com.twitter.util.jackson; + +import java.lang.annotation.Documented; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +public @interface NoValidatorImplConstraint { + + /** Every constraint annotation must define a message element of type String. */ + String message() default ""; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. + */ + Class[] payload() default {}; +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperFunctions.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperFunctions.scala new file mode 100644 index 0000000000..d4ed02fb1a --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperFunctions.scala @@ -0,0 +1,95 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.core.`type`.TypeReference +import com.fasterxml.jackson.databind.{ + JsonNode, + ObjectMapperCopier, + ObjectMapper => JacksonObjectMapper +} +import com.twitter.util.jackson.caseclass.exceptions.{ + CaseClassFieldMappingException, + CaseClassMappingException +} +import com.twitter.util.logging.Logging +import java.lang.reflect.{ParameterizedType, Type} +import org.scalatest.matchers.should.Matchers + +trait ScalaObjectMapperFunctions extends Matchers with Logging { + /* Abstract */ + protected def mapper: ScalaObjectMapper + + protected def parse[T: Manifest](string: String): T = + mapper.parse[T](string) + + protected def parse[T: Manifest](jsonNode: JsonNode): T = + mapper.parse[T](jsonNode) + + protected def parse[T: Manifest](obj: T): T = + mapper.parse[T](mapper.writeValueAsBytes(obj)) + + protected def generate(any: Any): String = + mapper.writeValueAsString(any) + + protected def assertJson[T: Manifest](obj: T, expected: String): Unit = { + val json = generate(obj) + JsonDiff.assertDiff(expected, json) + parse[T](json) should equal(obj) + } + + protected def copy(underlying: JacksonScalaObjectMapperType): JacksonScalaObjectMapperType = + ObjectMapperCopier.copy(underlying) + + protected def typeFromManifest(m: Manifest[_]): Type = + if (m.typeArguments.isEmpty) { + m.runtimeClass + } else { + new ParameterizedType { + override def getRawType: Class[_] = m.runtimeClass + override def getActualTypeArguments: Array[Type] = + m.typeArguments.map(typeFromManifest).toArray + override def getOwnerType: Null = null + } + } + + protected[this] def typeReference[T: Manifest]: TypeReference[T] = new TypeReference[T] { + override def getType: Type = typeFromManifest(manifest[T]) + } + + protected def deserialize[T: Manifest](mapper: JacksonObjectMapper, value: String): T = + mapper.readValue(value, typeReference[T]) + + protected def clearStackTrace( + exceptions: Seq[CaseClassFieldMappingException] + ): Seq[CaseClassFieldMappingException] = { + exceptions.foreach(_.setStackTrace(Array())) + exceptions + } + + protected def assertJsonParse[T: Manifest](json: String, withErrors: Seq[String]): Unit = { + if (withErrors.nonEmpty) { + val e1 = intercept[CaseClassMappingException] { + val parsed = parse[T](json) + println("Incorrectly parsed: " + ScalaObjectMapper().writePrettyString(parsed)) + } + assertObjectParseException(e1, withErrors) + + // also check that we can parse into an intermediate JsonNode + val e2 = intercept[CaseClassMappingException] { + val jsonNode = parse[JsonNode](json) + parse[T](jsonNode) + } + assertObjectParseException(e2, withErrors) + } else parse[T](json) + } + + private def assertObjectParseException( + e: CaseClassMappingException, + withErrors: Seq[String] + ): Unit = { + trace(e.errors.mkString("\n")) + clearStackTrace(e.errors) + + val actualMessages = e.errors.map(_.getMessage) + JsonDiff.assertDiff(withErrors, actualMessages) + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperTest.scala new file mode 100644 index 0000000000..ace3b4cbf6 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/ScalaObjectMapperTest.scala @@ -0,0 +1,2332 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.core.JsonFactory +import com.fasterxml.jackson.core.JsonParser +import com.fasterxml.jackson.databind.DeserializationContext +import com.fasterxml.jackson.databind.DeserializationFeature +import com.fasterxml.jackson.databind.JsonDeserializer +import com.fasterxml.jackson.databind.JsonMappingException +import com.fasterxml.jackson.databind.JsonNode +import com.fasterxml.jackson.databind.MapperFeature +import com.fasterxml.jackson.databind.PropertyNamingStrategies +import com.fasterxml.jackson.databind.module.SimpleModule +import com.fasterxml.jackson.databind.node.IntNode +import com.fasterxml.jackson.databind.node.TreeTraversingParser +import com.fasterxml.jackson.databind.{ObjectMapper => JacksonObjectMapper} +import com.fasterxml.jackson.dataformat.yaml.YAMLFactory +import com.fasterxml.jackson.module.scala.DefaultScalaModule +import com.fasterxml.jackson.module.scala.{ScalaObjectMapper => JacksonScalaObjectMapper} +import com.twitter.io.Buf +import com.twitter.util.Duration +import com.twitter.util.jackson.Obj.NestedCaseClassInObject +import com.twitter.util.jackson.Obj.NestedCaseClassInObjectWithNestedCaseClassInObjectParam +import com.twitter.util.jackson.TypeAndCompanion.NestedCaseClassInCompanion +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.Unspecified +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.ValidationError +import com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException +import com.twitter.util.jackson.internal.SimplePersonInPackageObject +import com.twitter.util.jackson.internal.SimplePersonInPackageObjectWithoutConstructorParams +import java.io.ByteArrayInputStream +import java.io.ByteArrayOutputStream +import java.time.LocalDateTime +import java.time.ZoneId +import java.time.format.DateTimeFormatter +import java.util.TimeZone +import java.util.concurrent.TimeUnit +import org.junit.Assert.assertNotNull +import org.junit.runner.RunWith +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import scala.util.Random + +private object ScalaObjectMapperTest { + val UTCZoneId: ZoneId = ZoneId.of("UTC") + val ISO8601DateTimeFormatter: DateTimeFormatter = + DateTimeFormatter.ISO_OFFSET_DATE_TIME // "2014-05-30T03:57:59.302Z" + + class ZeroOrOneDeserializer extends JsonDeserializer[ZeroOrOne] { + override def deserialize( + jsonParser: JsonParser, + deserializationContext: DeserializationContext + ): ZeroOrOne = { + jsonParser.getValueAsString match { + case "zero" => Zero + case "one" => One + case _ => + throw new JsonMappingException(null, "Invalid value") + } + } + } + + object MixInAnnotationsModule extends SimpleModule { + setMixInAnnotation(classOf[Point], classOf[PointMixin]) + setMixInAnnotation(classOf[CaseClassShouldUseKebabCaseFromMixin], classOf[KebabCaseMixin]) + } +} + +@RunWith(classOf[JUnitRunner]) +class ScalaObjectMapperTest + extends AnyFunSuite + with BeforeAndAfterAll + with Matchers + with ScalaObjectMapperFunctions { + + import ScalaObjectMapperTest._ + + /* Class under test */ + protected val mapper: ScalaObjectMapper = ScalaObjectMapper() + + private final val steve: Person = + Person(id = 1, name = "Steve", age = Some(20), age_with_default = Some(20), nickname = "ace") + private final val steveJson: String = + """{ + "id" : 1, + "name" : "Steve", + "age" : 20, + "age_with_default" : 20, + "nickname" : "ace" + } + """ + + override def beforeAll(): Unit = { + super.beforeAll() + mapper.registerModule(ScalaObjectMapperTest.MixInAnnotationsModule) + } + + test("constructors") { + // new mapper with defaults + assertNotNull(ScalaObjectMapper()) + assertNotNull(ScalaObjectMapper.apply()) + + // augments the given mapper with the defaults + assertNotNull(ScalaObjectMapper(mapper.underlying)) + assertNotNull(ScalaObjectMapper.apply(mapper.underlying)) + + // copies the underlying Jackson object mapper + assertNotNull(ScalaObjectMapper.objectMapper(mapper.underlying)) + // underlying mapper needs to have a JsonFactory of type YAMLFactory + assertThrows[IllegalArgumentException](ScalaObjectMapper.yamlObjectMapper(mapper.underlying)) + assertNotNull( + ScalaObjectMapper.yamlObjectMapper(ScalaObjectMapper.builder.yamlObjectMapper.underlying)) + assertNotNull(ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying)) + assertNotNull(ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying)) + + assertThrows[AssertionError](new ScalaObjectMapper(null)) + } + + test("builder constructors") { + assertNotNull(ScalaObjectMapper.builder.objectMapper) + assertNotNull(ScalaObjectMapper.builder.objectMapper(new JsonFactory)) + assertNotNull(ScalaObjectMapper.builder.objectMapper(new YAMLFactory)) + + assertNotNull(ScalaObjectMapper.builder.yamlObjectMapper) + assertNotNull(ScalaObjectMapper.builder.camelCaseObjectMapper) + assertNotNull(ScalaObjectMapper.builder.snakeCaseObjectMapper) + } + + test("mapper register module") { + val testMapper = ScalaObjectMapper() + + val simpleJacksonModule = new SimpleModule() + simpleJacksonModule.addDeserializer(classOf[ZeroOrOne], new ZeroOrOneDeserializer) + testMapper.registerModule(simpleJacksonModule) + + // regular mapper (without the custom deserializer) -- doesn't parse + intercept[JsonMappingException] { + mapper.parse[CaseClassWithZeroOrOne]("{\"id\" :\"zero\"}") + } + + // mapper with custom deserializer -- parses correctly + testMapper.parse[CaseClassWithZeroOrOne]("{\"id\" :\"zero\"}") should be( + CaseClassWithZeroOrOne(Zero)) + testMapper.parse[CaseClassWithZeroOrOne]("{\"id\" :\"one\"}") should be( + CaseClassWithZeroOrOne(One)) + intercept[JsonMappingException] { + testMapper.parse[CaseClassWithZeroOrOne]("{\"id\" :\"two\"}") + } + } + + test("inject request field fails with a mapping exception") { + // with no InjectableValues we get a CaseClassMappingException since we can't map + // the JSON into the case class. + intercept[CaseClassMappingException] { + parse[ClassWithQueryParamDateTimeInject]("""{}""") + } + } + + test( + "class with an injectable field fails with a mapping exception when it cannot be parsed from JSON") { + // if there is no injector, the default injectable values is not configured, thus the code + // tries to construct the asked for type from the given JSON, which is empty and fails with + // a mapping exception. + intercept[CaseClassMappingException] { + parse[ClassWithFooClassInject]("""{}""") + } + } + + test("class with an injectable field is not constructed from JSON") { + // no field injection configured and parsing proceeds as normal + val result = parse[ClassWithFooClassInject]( + """ + |{ + | "foo_class": { + | "id" : "11" + | } + |}""".stripMargin + ) + result.fooClass should equal(FooClass("11")) + } + + test("class with a defaulted injectable field is constructed from the default") { + // no field injection configured and default value is not used since a value is passed in the JSON. + val result = parse[ClassWithFooClassInjectAndDefault]( + """ + |{ + | "foo_class": { + | "id" : "1" + | } + |}""".stripMargin + ) + + result.fooClass should equal(FooClass("1")) + } + + test("injectable field should use default when field not sent in json") { + parse[CaseClassInjectStringWithDefault]("""{}""") should equal( + CaseClassInjectStringWithDefault("DefaultHello")) + } + + test("injectable field should not use default when field sent in json") { + parse[CaseClassInjectStringWithDefault]("""{"string": "123"}""") should equal( + CaseClassInjectStringWithDefault("123")) + } + + test("injectable field should use None assumed default when field not sent in json") { + parse[CaseClassInjectOptionString]("""{}""") should equal(CaseClassInjectOptionString(None)) + } + + test("injectable field should throw exception when no value is passed in the json") { + intercept[CaseClassMappingException] { + parse[CaseClassInjectString]("""{}""") + } + } + + test("regular mapper handles unknown properties when json provides MORE fields than case class") { + // regular mapper -- doesn't fail + mapper.parse[CaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + + // mapper = loose, case class = annotated strict --> Fail + intercept[JsonMappingException] { + mapper.parse[StrictCaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + } + } + + test("regular mapper handles unknown properties when json provides LESS fields than case class") { + // regular mapper -- doesn't fail + mapper.parse[CaseClassWithOption]( + """ + |{ + | "value": 12345, + | "extra": "fail" + |} + |""".stripMargin + ) + + // mapper = loose, case class = annotated strict --> Fail + intercept[JsonMappingException] { + mapper.parse[StrictCaseClassWithOption]( + """ + |{ + | "value": 12345, + | "extra": "fail" + |} + |""".stripMargin + ) + } + } + + test("regular mapper handles unknown properties") { + // regular mapper -- doesn't fail (no field named 'flame') + mapper.parse[CaseClassIdAndOption]( + """ + |{ + | "id": 12345, + | "flame": "gadget" + |} + |""".stripMargin + ) + + // mapper = loose, case class = annotated strict --> Fail (no field named 'flame') + intercept[JsonMappingException] { + mapper.parse[StrictCaseClassIdAndOption]( + """ + |{ + | "id": 12345, + | "flame": "gadget" + |} + |""".stripMargin + ) + } + } + + test("mapper with deserialization config fails on unknown properties") { + val testMapper = + ScalaObjectMapper.builder + .withDeserializationConfig(Map(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES -> true)) + .objectMapper + + // mapper = strict, case class = unannotated --> Fail + intercept[JsonMappingException] { + testMapper.parse[CaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + } + + // mapper = strict, case class = annotated strict --> Fail + intercept[JsonMappingException] { + testMapper.parse[StrictCaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + } + + // mapper = strict, case class = annotated loose --> Parse + testMapper.parse[LooseCaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "pass" + |} + |""".stripMargin + ) + } + + test("mapper with additional configuration handles unknown properties") { + // test with additional configuration set on mapper + val testMapper = + ScalaObjectMapper.builder + .withAdditionalMapperConfigurationFn( + _.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, true)) + .objectMapper + + // mapper = strict, case class = unannotated --> Fail + intercept[JsonMappingException] { + testMapper.parse[CaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + } + + // mapper = strict, case class = annotated strict --> Fail + intercept[JsonMappingException] { + testMapper.parse[StrictCaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "fail" + |} + |""".stripMargin + ) + } + + // mapper = strict, case class = annotated loose --> Parse + testMapper.parse[LooseCaseClass]( + """ + |{ + | "id": 12345, + | "name": "gadget", + | "extra": "pass" + |} + |""".stripMargin + ) + } + + test("support wrapped object mapper") { + val person = CamelCaseSimplePersonNoAnnotation(myName = "Bob") + + val jacksonScalaObjectMapper: JacksonScalaObjectMapperType = new JacksonObjectMapper() + with JacksonScalaObjectMapper + jacksonScalaObjectMapper.registerModule(DefaultScalaModule) + jacksonScalaObjectMapper.setPropertyNamingStrategy(PropertyNamingStrategies.LOWER_CAMEL_CASE) + + // default for ScalaObjectMapper#Builder is snake_case which should not get applied here + // since this method only wraps and does not mutate the underlying mapper + val objectMapper: ScalaObjectMapper = ScalaObjectMapper.objectMapper(jacksonScalaObjectMapper) + + val serialized = """{"myName":"Bob"}""" + // serialization -- they should each serialize the same way + jacksonScalaObjectMapper.writeValueAsString(person) should equal(serialized) + objectMapper.writeValueAsString(person) should equal(serialized) + jacksonScalaObjectMapper.writeValueAsString(person) should equal( + objectMapper.writeValueAsString(person)) + // deserialization -- each should be able to deserialize the other's written representation + objectMapper.parse[CamelCaseSimplePersonNoAnnotation]( + jacksonScalaObjectMapper.writeValueAsString(person)) should equal(person) + jacksonScalaObjectMapper.readValue[CamelCaseSimplePersonNoAnnotation]( + objectMapper.writeValueAsString(person)) should equal(person) + objectMapper.parse[CamelCaseSimplePersonNoAnnotation](serialized) should equal( + jacksonScalaObjectMapper.readValue[CamelCaseSimplePersonNoAnnotation](serialized)) + } + + test("support camel case mapper") { + val camelCaseObjectMapper = ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying) + + camelCaseObjectMapper.parse[Map[String, String]]("""{"firstName": "Bob"}""") should equal( + Map("firstName" -> "Bob") + ) + } + + test("support snake case mapper") { + val snakeCaseObjectMapper = ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying) + + val person = CamelCaseSimplePersonNoAnnotation(myName = "Bob") + + val serialized = snakeCaseObjectMapper.writeValueAsString(person) + serialized should equal("""{"my_name":"Bob"}""") + snakeCaseObjectMapper.parse[CamelCaseSimplePersonNoAnnotation](serialized) should equal(person) + } + + test("support yaml mapper") { + val yamlObjectMapper = ScalaObjectMapper.builder.yamlObjectMapper + + val person = CamelCaseSimplePersonNoAnnotation(myName = "Bob") + + val serialized = yamlObjectMapper.writeValueAsString(person) + // default PropertyNamingStrategy for the generate YAML object mapper is snake_case + serialized should equal("""--- + |my_name: "Bob" + |""".stripMargin) + yamlObjectMapper.parse[CamelCaseSimplePersonNoAnnotation](serialized) should equal(person) + + // but we can also update to a camelCase PropertyNamingStrategy: + val camelCaseObjectMapper = ScalaObjectMapper.camelCaseObjectMapper(yamlObjectMapper.underlying) + val serializedCamelCase = camelCaseObjectMapper.writeValueAsString(person) + serializedCamelCase should equal("""--- + |myName: "Bob" + |""".stripMargin) + camelCaseObjectMapper.parse[CamelCaseSimplePersonNoAnnotation]( + serializedCamelCase) should equal(person) + } + + test("Scala enums") { + val expectedInstance = BasicDate(Month.Feb, 29, 2020, Weekday.Sat) + val expectedStr = """{"month":"Feb","day":29,"year":2020,"weekday":"Sat"}""" + parse[BasicDate](expectedStr) should equal(expectedInstance) // deser + generate(expectedInstance) should equal(expectedStr) // ser + + // test with default scala module + val defaultScalaObjectMapper = new JacksonObjectMapper with JacksonScalaObjectMapper + defaultScalaObjectMapper.registerModule(DefaultScalaModule) + defaultScalaObjectMapper.setPropertyNamingStrategy(PropertyNamingStrategies.SNAKE_CASE) + deserialize[BasicDate](defaultScalaObjectMapper, expectedStr) should equal(expectedInstance) + + // case insensitive mapper feature does not pertain to scala enumerations (only Java enums) + val caseInsensitiveEnumMapper = + ScalaObjectMapper.builder + .withAdditionalMapperConfigurationFn( + _.configure(MapperFeature.ACCEPT_CASE_INSENSITIVE_ENUMS, true)) + .objectMapper + + val withErrors212 = Seq("month: key not found: feb", "weekday: key not found: sat") + val withErrors213 = Seq("month: null", "weekday: null") + + val expectedCaseInsensitiveStr = """{"month":"feb","day":29,"year":2020,"weekday":"sat"}""" + val e = intercept[CaseClassMappingException] { + caseInsensitiveEnumMapper.parse[BasicDate](expectedCaseInsensitiveStr) + } + val actualMessages = e.errors.map(_.getMessage) + assert( + actualMessages.mkString(",") == withErrors212.mkString(",") || + actualMessages.mkString(",") == withErrors213.mkString(",")) + + // non-existent values + val withErrors2212 = Seq("month: key not found: Nnn", "weekday: key not found: Flub") + val withErrors2213 = Seq("month: null", "weekday: null") + + val expectedStr2 = """{"month":"Nnn","day":29,"year":2020,"weekday":"Flub"}""" + val e2 = intercept[CaseClassMappingException] { + parse[BasicDate](expectedStr2) + } + val actualMessages2 = e2.errors.map(_.getMessage) + assert( + actualMessages2.mkString(",") == withErrors2212.mkString(",") || + actualMessages2.mkString(",") == withErrors2213.mkString(",")) + + val invalidWithOptionalScalaEnumerationValue = """{"month":"Nnn"}""" + val e3 = intercept[CaseClassMappingException] { + parse[WithOptionalScalaEnumeration](invalidWithOptionalScalaEnumerationValue) + } + e3.errors.size shouldEqual 1 + e3.errors.head.getMessage match { + case "month: key not found: Nnn" => // Scala 2.12.x (https://github.com/scala/scala/blob/2.12.x/src/library/scala/collection/MapLike.scala#L236) + case "month: null" => // Scala 2.13.1 (https://github.com/scala/scala/blob/2.13.x/src/library/scala/collection/immutable/HashMap.scala#L635) + case _ => fail() + } + e3.errors.head.reason.message match { + case "key not found: Nnn" => // Scala 2.12.x (https://github.com/scala/scala/blob/2.12.x/src/library/scala/collection/MapLike.scala#L236_ + case null => // Scala 2.13.1 (https://github.com/scala/scala/blob/2.13.x/src/library/scala/collection/immutable/HashMap.scala#L635) + case _ => fail() + } + e3.errors.head.reason.detail shouldEqual Unspecified + + val validWithOptionalScalaEnumerationValue = """{"month":"Apr"}""" + val validWithOptionalScalaEnumeration = + parse[WithOptionalScalaEnumeration](validWithOptionalScalaEnumerationValue) + validWithOptionalScalaEnumeration.month.isDefined shouldBe true + validWithOptionalScalaEnumeration.month.get shouldEqual Month.Apr + } + + // based on Jackson test to ensure compatibility: + // https://github.com/FasterXML/jackson-module-scala/blob/fa7cf702e0f61467d726384af88de9ea1f798b97/src/test/scala/com/fasterxml/jackson/module/scala/deser/CaseClassDeserializerTest.scala#L78-L81 + test("deserialization#generic types") { + parse[GenericTestCaseClass[Int]]("""{"data" : 3}""") should equal(GenericTestCaseClass(3)) + parse[GenericTestCaseClass[CaseClass]]("""{"data" : {"id" : 123, "name" : "foo"}}""") should + equal(GenericTestCaseClass(CaseClass(123, "foo"))) + parse[CaseClassWithGeneric[CaseClass]]( + """{"inside" : { "data" : {"id" : 123, "name" : "foo"}}}""") should + equal(CaseClassWithGeneric(GenericTestCaseClass(CaseClass(123, "foo")))) + + parse[GenericTestCaseClass[String]]("""{"data" : "Hello, World"}""") should equal( + GenericTestCaseClass("Hello, World")) + parse[GenericTestCaseClass[Double]]("""{"data" : 3.14}""") should equal( + GenericTestCaseClass(3.14d)) + + parse[CaseClassWithOptionalGeneric[Int]]("""{"inside": {"data": 3}}""") should equal( + CaseClassWithOptionalGeneric[Int](Some(GenericTestCaseClass[Int](3)))) + + parse[CaseClassWithTypes[String, Int]]("""{"first": "Bob", "second" : 42}""") should equal( + CaseClassWithTypes("Bob", 42)) + parse[CaseClassWithTypes[Int, Float]]("""{"first": 127, "second" : 39.0}""") should equal( + CaseClassWithTypes(127, 39.0f)) + + parse[CaseClassWithMapTypes[String, Float]]( + """{"data": {"pi": 3.14, "inverse fine structure constant": 137.035}}""") should equal( + CaseClassWithMapTypes( + Map[String, Float]("pi" -> 3.14f, "inverse fine structure constant" -> 137.035f)) + ) + parse[CaseClassWithManyTypes[Int, Float, String]]( + """{"one": 1, "two": 3.1, "three": "Hello, World!"}""") should equal( + CaseClassWithManyTypes(1, 3.1f, "Hello, World!")) + + val result = Page( + List( + Person(1, "Bob Marley", None, Some(32), Some(32), "Music Master"), + Person(2, "Jimi Hendrix", None, Some(27), None, "Melody Man")), + 5, + None, + None) + val input = + """ + |{ + | "data": [ + | {"id": 1, "name": "Bob Marley", "age": 32, "age_with_default": 32, "nickname": "Music Master"}, + | {"id": 2, "name": "Jimi Hendrix", "age": 27, "nickname": "Melody Man"} + | ], + | "page_size": 5 + |} + """.stripMargin + + // test with default scala module + val defaultScalaObjectMapper = new JacksonObjectMapper with JacksonScalaObjectMapper + defaultScalaObjectMapper.registerModule(DefaultScalaModule) + defaultScalaObjectMapper.setPropertyNamingStrategy(PropertyNamingStrategies.SNAKE_CASE) + deserialize[Page[Person]](defaultScalaObjectMapper, input) should equal(result) + + // test with ScalaObjectMapper + deserialize[Page[Person]](mapper.underlying, input) should equal(result) + } + + // tests for GH #547 + test("deserialization#generic types 2") { + val aGeneric = mapper.parse[AGeneric[B]]("""{"b":{"value":"string"}}""") + aGeneric.b.map(_.value) should equal(Some("string")) + + val c = mapper.parse[C]("""{"a":{"b":{"value":"string"}}}""") + c.a.b.map(_.value) should equal(Some("string")) + + val aaGeneric = mapper.parse[AAGeneric[B, Int]]("""{"b":{"value":"string"},"c":42}""") + aaGeneric.b.map(_.value) should equal(Some("string")) + aaGeneric.c.get should equal(42) + + val aaGeneric2 = mapper.parse[AAGeneric[B, D]]("""{"b":{"value":"string"},"c":{"value":42}}""") + aaGeneric2.b.map(_.value) should equal(Some("string")) + aaGeneric2.c.map(_.value) should equal(Some(42)) + + val eGeneric = + mapper.parse[E[B, D, Int]]("""{"a": {"b":{"value":"string"},"c":{"value":42}}, "b":42}""") + eGeneric.a.b.map(_.value) should equal(Some("string")) + eGeneric.a.c.map(_.value) should equal(Some(42)) + eGeneric.b should equal(Some(42)) + + val fGeneric = + mapper.parse[F[B, Int, D, AGeneric[B], Int, String]](""" + |{ + | "a":{"value":"string"}, + | "b":[1,2,3], + | "c":{"value":42}, + | "d":{"b":{"value":"string"}}, + | "e":{"r":"forty-two"} + |} + |""".stripMargin) + fGeneric.a.map(_.value) should equal(Some("string")) + fGeneric.b should equal(Seq(1, 2, 3)) + fGeneric.c.map(_.value) should equal(Some(42)) + fGeneric.d.b.map(_.value) should equal(Some("string")) + fGeneric.e should equal(Right("forty-two")) + + val fGeneric2 = + mapper.parse[F[B, Int, D, AGeneric[B], Int, String]](""" + |{ + | "a":{"value":"string"}, + | "b":[1,2,3], + | "c":{"value":42}, + | "d":{"b":{"value":"string"}}, + | "e":{"l":42} + |} + |""".stripMargin) + fGeneric2.a.map(_.value) should equal(Some("string")) + fGeneric2.b should equal(Seq(1, 2, 3)) + fGeneric2.c.map(_.value) should equal(Some(42)) + fGeneric2.d.b.map(_.value) should equal(Some("string")) + fGeneric2.e should equal(Left(42)) + + // requires polymorphic handling with @JsonTypeInfo + val withTypeBounds = + mapper.parse[WithTypeBounds[BoundaryA]](""" + |{ + | "a": {"type":"a", "value":"Guineafowl"} + |} + |""".stripMargin) + withTypeBounds.a.map(_.value) should equal(Some("Guineafowl")) + + // uses a specific json deserializer + val someNumberType = mapper.parse[SomeNumberType[java.lang.Integer]]( + """ + |{ + | "n": 42 + |} + |""".stripMargin + ) + someNumberType.n should equal(Some(42)) + + // multi-parameter non-collection type + val gGeneric = + mapper.parse[G[B, Int, D, AGeneric[B], Int, String]](""" + |{ + | "gg": { + | "t":{"value":"string"}, + | "u":42, + | "v":{"value":42}, + | "x":{"b":{"value":"string"}}, + | "y":137, + | "z":"string" + | } + |} + |""".stripMargin) + gGeneric.gg.t.value should be("string") + gGeneric.gg.u should equal(42) + gGeneric.gg.v.value should equal(42) + gGeneric.gg.x.b.map(_.value) should equal(Some("string")) + gGeneric.gg.y should equal(137) + gGeneric.gg.z should be("string") + } + + test("JsonProperty#annotation inheritance") { + val aumJson = """{"i":1,"j":"J"}""" + val aum = parse[Aum](aumJson) + aum should equal(Aum(1, "J")) + mapper.writeValueAsString(Aum(1, "J")) should equal(aumJson) + + val testCaseClassJson = """{"fedoras":["felt","straw"],"oldness":27}""" + val testCaseClass = parse[CaseClassTraitImpl](testCaseClassJson) + testCaseClass should equal(CaseClassTraitImpl(Seq("felt", "straw"), 27)) + mapper.writeValueAsString(CaseClassTraitImpl(Seq("felt", "straw"), 27)) should equal( + testCaseClassJson) + } + + test("simple tests#parse super simple") { + val foo = parse[SimplePerson]("""{"name": "Steve"}""") + foo should equal(SimplePerson("Steve")) + } + + test("simple tests#parse super simple -- YAML") { + val yamlMapper = ScalaObjectMapper.builder.yamlObjectMapper + val foo = yamlMapper.parse[SimplePerson]("""name: Steve""") + foo should equal(SimplePerson("Steve")) + + yamlMapper.writeValueAsString(foo) should equal("""--- + |name: "Steve" + |""".stripMargin) + } + + test("get PropertyNamingStrategy") { + val namingStrategy = mapper.propertyNamingStrategy + namingStrategy should not be null + } + + test("simple tests#parse simple") { + val foo = parse[SimplePerson]("""{"name": "Steve"}""") + foo should equal(SimplePerson("Steve")) + } + + test("simple tests#parse CamelCase simple person") { + val foo = parse[CamelCaseSimplePerson]("""{"myName": "Steve"}""") + foo should equal(CamelCaseSimplePerson("Steve")) + } + + test("simple tests#parse json") { + val person = parse[Person](steveJson) + person should equal(steve) + } + + test("simple tests#parse json list of objects") { + val json = Seq(steveJson, steveJson).mkString("[", ", ", "]") + val persons = parse[Seq[Person]](json) + persons should equal(Seq(steve, steve)) + } + + test("simple tests#parse json list of ints") { + val nums = parse[Seq[Int]]("""[1,2,3]""") + nums should equal(Seq(1, 2, 3)) + } + + test("simple tests#serialize case class with logging") { + val steveWithLogging = PersonWithLogging( + id = 1, + name = "Steve", + age = Some(20), + age_with_default = Some(20), + nickname = "ace" + ) + + assertJson(steveWithLogging, steveJson) + } + + test("simple tests#parse json with extra field at end") { + val person = parse[Person](""" + { + "id" : 1, + "name" : "Steve", + "age" : 20, + "age_with_default" : 20, + "nickname" : "ace", + "extra" : "extra" + } + """) + person should equal(steve) + } + + test("simple tests#parse json with extra field in middle") { + val person = parse[Person](""" + { + "id" : 1, + "name" : "Steve", + "age" : 20, + "extra" : "extra", + "age_with_default" : 20, + "nickname" : "ace" + } + """) + person should equal(steve) + } + + test("simple tests#parse json with extra field name with dot") { + val person = parse[PersonWithDottedName](""" + { + "id" : 1, + "name.last" : "Cosenza" + } + """) + + person should equal( + PersonWithDottedName( + id = 1, + lastName = "Cosenza" + ) + ) + } + + test("simple tests#parse json with missing 'id' and 'name' field and invalid age field") { + assertJsonParse[Person]( + """ { + "age" : "foo", + "age_with_default" : 20, + "nickname" : "ace" + }""", + withErrors = + Seq("age: 'foo' is not a valid Integer", "id: field is required", "name: field is required") + ) + } + + test("simple tests#parse nested json with missing fields") { + assertJsonParse[Car]( + """ + { + "id" : 0, + "make": "Foo", + "year": 2000, + "passengers" : [ { "id": "-1", "age": "blah" } ] + } + """, + withErrors = Seq( + "make: 'Foo' is not a valid CarMake with valid values: Ford, Honda", + "model: field is required", + "owners: field is required", + "passengers.age: 'blah' is not a valid Integer", + "passengers.name: field is required" + ) + ) + } + + test("simple test#parse char") { + assertJsonParse[CaseClassCharacter]( + """ + |{ + |"c" : -1 + |} + """.stripMargin, + withErrors = Seq( + "c: '-1' is not a valid Character" + ) + ) + } + + test("simple tests#parse json with missing 'nickname' field that has a string default") { + val person = parse[Person](""" + { + "id" : 1, + "name" : "Steve", + "age" : 20, + "age_with_default" : 20 + }""") + person should equal(steve.copy(nickname = "unknown")) + } + + test( + "simple tests#parse json with missing 'age' field that is an Option without a default should succeed" + ) { + parse[Person](""" + { + "id" : 1, + "name" : "Steve", + "age_with_default" : 20, + "nickname" : "bob" + } + """) + } + + test("simple tests#parse json into JsonNode") { + parse[JsonNode](steveJson) + } + + test("simple tests#generate json") { + assertJson(steve, steveJson) + } + + test("simple tests#generate then parse") { + val json = generate(steve) + val person = parse[Person](json) + person should equal(steve) + } + + test("simple tests#generate then parse Either type") { + type T = Either[String, Int] + val l: T = Left("Q?") + val r: T = Right(42) + assertJson(l, """{"l":"Q?"}""") + assertJson(r, """{"r":42}""") + } + + test("simple tests#generate then parse nested case class") { + val origCar = Car(1, CarMake.Ford, "Explorer", 2001, Seq(steve, steve)) + val carJson = generate(origCar) + val car = parse[Car](carJson) + car should equal(origCar) + } + + test("simple tests#Prevent overwriting val in case class") { + parse[CaseClassWithVal]("""{ + "name" : "Bob", + "type" : "dog" + }""") should equal(CaseClassWithVal("Bob")) + } + + test("simple tests#parse WithEmptyJsonProperty then write and see if it equals original") { + val withEmptyJsonProperty = + """{ + | "foo" : "abc" + |}""".stripMargin + val obj = parse[WithEmptyJsonProperty](withEmptyJsonProperty) + val json = mapper.writePrettyString(obj) + json should equal(withEmptyJsonProperty) + } + + test("simple tests#parse WithNonemptyJsonProperty then write and see if it equals original") { + val withNonemptyJsonProperty = + """{ + | "bar" : "abc" + |}""".stripMargin + val obj = parse[WithNonemptyJsonProperty](withNonemptyJsonProperty) + val json = mapper.writePrettyString(obj) + json should equal(withNonemptyJsonProperty) + } + + test( + "simple tests#parse WithoutJsonPropertyAnnotation then write and see if it equals original") { + val withoutJsonPropertyAnnotation = + """{ + | "foo" : "abc" + |}""".stripMargin + val obj = parse[WithoutJsonPropertyAnnotation](withoutJsonPropertyAnnotation) + val json = mapper.writePrettyString(obj) + json should equal(withoutJsonPropertyAnnotation) + } + + test( + "simple tests#use default Jackson mapper without setting naming strategy to see if it remains camelCase to verify default Jackson behavior" + ) { + val objMapper = new JacksonObjectMapper with JacksonScalaObjectMapper + objMapper.registerModule(DefaultScalaModule) + + val response = objMapper.writeValueAsString(NamingStrategyJsonProperty("abc")) + response should equal("""{"longFieldName":"abc"}""") + } + + test( + "simple tests#use default Jackson mapper after setting naming strategy and see if it changes to verify default Jackson behavior" + ) { + val objMapper = new JacksonObjectMapper with JacksonScalaObjectMapper + objMapper.registerModule(DefaultScalaModule) + objMapper.setPropertyNamingStrategy(new PropertyNamingStrategies.SnakeCaseStrategy) + + val response = objMapper.writeValueAsString(NamingStrategyJsonProperty("abc")) + response should equal("""{"long_field_name":"abc"}""") + } + + test("enums#simple") { + parse[CaseClassWithEnum]("""{ + "name" : "Bob", + "make" : "ford" + }""") should equal(CaseClassWithEnum("Bob", CarMakeEnum.ford)) + } + + test("enums#complex") { + JsonDiff.assertDiff( + CaseClassWithComplexEnums( + "Bob", + CarMakeEnum.vw, + Some(CarMakeEnum.ford), + Seq(CarMakeEnum.vw, CarMakeEnum.ford), + Set(CarMakeEnum.ford, CarMakeEnum.vw) + ), + parse[CaseClassWithComplexEnums]("""{ + "name" : "Bob", + "make" : "vw", + "make_opt" : "ford", + "make_seq" : ["vw", "ford"], + "make_set" : ["ford", "vw"] + }""") + ) + } + + test("enums#complex case insensitive") { + val caseInsensitiveEnumMapper = + ScalaObjectMapper.builder + .withAdditionalMapperConfigurationFn( + _.configure(MapperFeature.ACCEPT_CASE_INSENSITIVE_ENUMS, true)) + .objectMapper + + JsonDiff.assertDiff( + CaseClassWithComplexEnums( + "Bob", + CarMakeEnum.vw, + Some(CarMakeEnum.ford), + Seq(CarMakeEnum.vw, CarMakeEnum.ford), + Set(CarMakeEnum.ford, CarMakeEnum.vw) + ), + caseInsensitiveEnumMapper.parse[CaseClassWithComplexEnums]("""{ + "name" : "Bob", + "make" : "VW", + "make_opt" : "Ford", + "make_seq" : ["vW", "foRD"], + "make_set" : ["fORd", "Vw"] + }""") + ) + } + + test("enums#invalid enum entry") { + val e = intercept[CaseClassMappingException] { + parse[CaseClassWithEnum]("""{ + "name" : "Bob", + "make" : "foo" + }""") + } + e.errors.map { _.getMessage } should equal( + Seq("""make: 'foo' is not a valid CarMakeEnum with valid values: ford, vw""") + ) + } + + test("LocalDateTime support") { + parse[CaseClassWithDateTime]("""{ + "date_time" : "2014-05-30T03:57:59.302Z" + }""") should equal( + CaseClassWithDateTime( + LocalDateTime.parse("2014-05-30T03:57:59.302Z", ISO8601DateTimeFormatter)) + ) + } + + test("invalid LocalDateTime 1") { + assertJsonParse[CaseClassWithDateTime]( + """{ + "date_time" : "" + }""", + withErrors = Seq("date_time: error parsing ''")) + } + + test("invalid LocalDateTime 2") { + assertJsonParse[CaseClassWithIntAndDateTime]( + """{ + "name" : "Bob", + "age" : "old", + "age2" : "1", + "age3" : "", + "date_time" : "today", + "date_time2" : "1", + "date_time3" : -1, + "date_time4" : "" + }""", + withErrors = Seq( + "age3: '' is not a valid Integer", + "age: 'old' is not a valid Integer", + "date_time2: '1' is not a valid LocalDateTime", + "date_time3: '' is not a valid LocalDateTime", + "date_time4: error parsing ''", + "date_time: 'today' is not a valid LocalDateTime" + ) + ) + } + + test("TwitterUtilDuration serialize") { + val serialized = mapper.writeValueAsString( + CaseClassWithTwitterUtilDuration(Duration.fromTimeUnit(3L, TimeUnit.HOURS)) + ) + serialized should equal("""{"duration":"3.hours"}""") + } + + test("TwitterUtilDuration deserialize") { + parse[CaseClassWithTwitterUtilDuration]("""{ + "duration": "3.hours" + }""") should equal( + CaseClassWithTwitterUtilDuration(Duration.fromTimeUnit(3L, TimeUnit.HOURS)) + ) + } + + test("TwitterUtilDuration deserialize invalid") { + assertJsonParse[CaseClassWithTwitterUtilDuration]( + """{"duration": "3.unknowns"}""", + withErrors = Seq("duration: Invalid unit: unknowns") + ) + } + + test("escaped fields#long") { + parse[CaseClassWithEscapedLong]("""{ + "1-5" : 10 + }""") should equal(CaseClassWithEscapedLong(`1-5` = 10)) + } + + test("escaped fields#string") { + parse[CaseClassWithEscapedString]("""{ + "1-5" : "10" + }""") should equal(CaseClassWithEscapedString(`1-5` = "10")) + } + + test("escaped fields#non-unicode escaped") { + parse[CaseClassWithEscapedNormalString]("""{ + "a" : "foo" + }""") should equal(CaseClassWithEscapedNormalString("foo")) + } + + test("escaped fields#unicode and non-unicode fields") { + parse[UnicodeNameCaseClass]("""{"winning-id":23,"name":"the name of this"}""") should equal( + UnicodeNameCaseClass(23, "the name of this") + ) + } + + test("wrapped values#support value class") { + val orig = WrapperValueClass(1) + val obj = parse[WrapperValueClass](generate(orig)) + orig should equal(obj) + } + + test("wrapped values#direct WrappedValue for Int") { + val origObj = WrappedValueInt(1) + val obj = parse[WrappedValueInt](generate(origObj)) + origObj should equal(obj) + } + + test("wrapped values#direct WrappedValue for String") { + val origObj = WrappedValueString("1") + val obj = parse[WrappedValueString](generate(origObj)) + origObj should equal(obj) + } + + test( + "wrapped values#direct WrappedValue for String when asked to parse wrapped json object should throw exception" + ) { + intercept[JsonMappingException] { + parse[WrappedValueString]("""{"value": "1"}""") + } + } + + test("wrapped values#direct WrappedValue for Long") { + val origObj = WrappedValueLong(1) + val obj = parse[WrappedValueLong](generate(origObj)) + origObj should equal(obj) + } + + test("wrapped values#WrappedValue for Int") { + val origObj = WrappedValueIntInObj(WrappedValueInt(1)) + val json = mapper.writeValueAsString(origObj) + val obj = parse[WrappedValueIntInObj](json) + origObj should equal(obj) + } + + test("wrapped values#WrappedValue for String") { + val origObj = WrappedValueStringInObj(WrappedValueString("1")) + val json = mapper.writeValueAsString(origObj) + val obj = parse[WrappedValueStringInObj](json) + origObj should equal(obj) + } + + test("wrapped values#WrappedValue for Long") { + val origObj = WrappedValueLongInObj(WrappedValueLong(11111111)) + val json = mapper.writeValueAsString(origObj) + val obj = parse[WrappedValueLongInObj](json) + origObj should equal(obj) + } + + test("wrapped values#Seq[WrappedValue]") { + generate(Seq(WrappedValueLong(11111111))) should be("""[11111111]""") + } + + test("wrapped values#Map[WrappedValueString, String]") { + val obj = Map(WrappedValueString("11111111") -> "asdf") + val json = generate(obj) + json should be("""{"11111111":"asdf"}""") + parse[Map[WrappedValueString, String]](json) should be(obj) + } + + test("wrapped values#Map[WrappedValueLong, String]") { + assertJson(Map(WrappedValueLong(11111111) -> "asdf"), """{"11111111":"asdf"}""") + } + + test("wrapped values#deser Map[Long, String]") { + val obj = parse[Map[Long, String]]("""{"11111111":"asdf"}""") + val expected = Map(11111111L -> "asdf") + obj should equal(expected) + } + + test("wrapped values#deser Map[String, String]") { + parse[Map[String, String]]("""{"11111111":"asdf"}""") should + be(Map("11111111" -> "asdf")) + } + + test("wrapped values#Map[String, WrappedValueLong]") { + generate(Map("asdf" -> WrappedValueLong(11111111))) should be("""{"asdf":11111111}""") + } + + test("wrapped values#Map[String, WrappedValueString]") { + generate(Map("asdf" -> WrappedValueString("11111111"))) should be("""{"asdf":"11111111"}""") + } + + test("wrapped values#object with Map[WrappedValueString, String]") { + assertJson( + obj = WrappedValueStringMapObject(Map(WrappedValueString("11111111") -> "asdf")), + expected = """{ + "map" : { + "11111111":"asdf" + } + }""" + ) + } + + test("fail when CaseClassWithSeqLongs with null array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqOfLongs]("""{"seq": [null]}""") + } + } + + test("fail when CaseClassWithSeqWrappedValueLong with null array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqWrappedValueLong]("""{"seq": [null]}""") + } + } + + test("fail when CaseClassWithArrayWrappedValueLong with null array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithArrayWrappedValueLong]("""{"array": [null]}""") + } + } + + test("fail when CaseClassWithArrayLong with null field in object") { + intercept[CaseClassMappingException] { + parse[CaseClassWithArrayLong]("""{"array": [null]}""") + } + } + + test("fail when CaseClassWithArrayBoolean with null field in object") { + intercept[CaseClassMappingException] { + parse[CaseClassWithArrayBoolean]("""{"array": [null]}""") + } + } + + // ======================================================== + // Jerkson Inspired/Copied Tests Below + test("A basic case class generates a JSON object with matching field values") { + generate(CaseClass(1, "Coda")) should be("""{"id":1,"name":"Coda"}""") + } + + test("A basic case class is parsable from a JSON object with corresponding fields") { + parse[CaseClass]("""{"id":111,"name":"Coda"}""") should be(CaseClass(111L, "Coda")) + } + + test("A basic case class is parsable from a JSON object with extra fields") { + parse[CaseClass]("""{"id":1,"name":"Coda","derp":100}""") should be(CaseClass(1, "Coda")) + } + + test("A basic case class is not parsable from an incomplete JSON object") { + intercept[Exception] { + parse[CaseClass]("""{"id":1}""") + } + } + + test("A case class with lazy fields generates a JSON object with those fields evaluated") { + generate(CaseClassWithLazyVal(1)) should be("""{"id":1,"woo":"yeah"}""") + } + + test("A case class with lazy fields is parsable from a JSON object without those fields") { + parse[CaseClassWithLazyVal]("""{"id":1}""") should be(CaseClassWithLazyVal(1)) + } + + test("A case class with lazy fields is not parsable from an incomplete JSON object") { + intercept[Exception] { + parse[CaseClassWithLazyVal]("""{}""") + } + } + + test("A case class with ignored members generates a JSON object without those fields") { + generate(CaseClassWithIgnoredField(1)) should be("""{"id":1}""") + generate(CaseClassWithIgnoredFieldsExactMatch(1)) should be("""{"id":1}""") + generate(CaseClassWithIgnoredFieldsMatchAfterToSnakeCase(1)) should be("""{"id":1}""") + } + + test("A case class with some ignored members is not parsable from an incomplete JSON object") { + intercept[Exception] { + parse[CaseClassWithIgnoredField]("""{}""") + } + intercept[Exception] { + parse[CaseClassWithIgnoredFieldsMatchAfterToSnakeCase]("""{}""") + } + } + + test("A case class with transient members generates a JSON object without those fields") { + generate(CaseClassWithTransientField(1)) should be("""{"id":1}""") + } + + test("A case class with transient members is parsable from a JSON object without those fields") { + parse[CaseClassWithTransientField]("""{"id":1}""") should be(CaseClassWithTransientField(1)) + } + + test("A case class with some transient members is not parsable from an incomplete JSON object") { + intercept[Exception] { + parse[CaseClassWithTransientField]("""{}""") + } + } + + test("A case class with lazy vals generates a JSON object without those fields") { + // jackson-module-scala serializes lazy vals and the only work around is to use @JsonIgnore on the field + // see: https://github.com/FasterXML/jackson-module-scala/issues/113 + generate(CaseClassWithLazyField(1)) should be("""{"id":1}""") + } + + test("A case class with lazy vals is parsable from a JSON object without those fields") { + parse[CaseClassWithLazyField]("""{"id":1}""") should be(CaseClassWithLazyField(1)) + } + + test( + "A case class with an overloaded field generates a JSON object with the nullary version of that field" + ) { + generate(CaseClassWithOverloadedField(1)) should be("""{"id":1}""") + } + + test("A case class with an Option[String] member generates a field if the member is Some") { + generate(CaseClassWithOption(Some("what"))) should be("""{"value":"what"}""") + } + + test( + "A case class with an Option[String] member is parsable from a JSON object with that field") { + parse[CaseClassWithOption]("""{"value":"what"}""") should be(CaseClassWithOption(Some("what"))) + } + + test( + "A case class with an Option[String] member doesn't generate a field if the member is None") { + generate(CaseClassWithOption(None)) should be("""{}""") + } + + test( + "A case class with an Option[String] member is parsable from a JSON object without that field" + ) { + parse[CaseClassWithOption]("""{}""") should be(CaseClassWithOption(None)) + } + + test( + "A case class with an Option[String] member is parsable from a JSON object with a null value for that field" + ) { + parse[CaseClassWithOption]("""{"value":null}""") should be(CaseClassWithOption(None)) + } + + test("A case class with a JsonNode member generates a field of the given type") { + generate(CaseClassWithJsonNode(new IntNode(2))) should be("""{"value":2}""") + } + + test("issues#standalone map") { + val map = parse[Map[String, String]](""" + { + "one": "two" + } + """) + map should equal(Map("one" -> "two")) + } + + test("issues#case class with map") { + val obj = parse[CaseClassWithMap](""" + { + "map": {"one": "two"} + } + """) + obj.map should equal(Map("one" -> "two")) + } + + test("issues#case class with multiple constructors") { + intercept[AssertionError] { + parse[CaseClassWithTwoConstructors]("""{"id":42,"name":"TwoConstructorsNoDefault"}"""") + } + } + + test("issues#case class with multiple (3) constructors") { + intercept[AssertionError] { + parse[CaseClassWithThreeConstructors]("""{"id":42,"name":"ThreeConstructorsNoDefault"}"""") + } + } + + test("issues#case class nested within an object") { + parse[NestedCaseClassInObject](""" + { + "id": "foo" + } + """) should equal(NestedCaseClassInObject(id = "foo")) + } + + test( + "issues#case class nested within an object with member that is also a case class in an object" + ) { + parse[NestedCaseClassInObjectWithNestedCaseClassInObjectParam]( + """ + { + "nested": { + "id": "foo" + } + } + """ + ) should equal( + NestedCaseClassInObjectWithNestedCaseClassInObjectParam( + nested = NestedCaseClassInObject(id = "foo") + ) + ) + } + + test("deserialization#case class nested within a companion object works") { + parse[NestedCaseClassInCompanion]( + """ + { + "id": "foo" + } + """ + ) should equal(NestedCaseClassInCompanion(id = "foo")) + } + + case class NestedCaseClassInClass(id: String) + test("deserialization#case class nested within a class fails") { + intercept[AssertionError] { + parse[NestedCaseClassInClass](""" + { + "id": "foo" + } + """) should equal(NestedCaseClassInClass(id = "foo")) + } + } + + test("deserialization#case class with set of longs") { + val obj = parse[CaseClassWithSetOfLongs](""" + { + "set": [5000000, 1, 2, 3, 1000] + } + """) + obj.set.toSeq.sorted should equal(Seq(1L, 2L, 3L, 1000L, 5000000L)) + } + + test("deserialization#case class with seq of longs") { + val obj = parse[CaseClassWithSeqOfLongs](""" + { + "seq": [%s] + } + """.format(1004 to 1500 mkString ",")) + obj.seq.sorted should equal(1004 to 1500) + } + + test("deserialization#nested case class with collection of longs") { + val idsStr = 1004 to 1500 mkString "," + val obj = parse[CaseClassWithNestedSeqLong](""" + { + "seq_class" : {"seq": [%s]}, + "set_class" : {"set": [%s]} + } + """.format(idsStr, idsStr)) + obj.seqClass.seq.sorted should equal(1004 to 1500) + obj.setClass.set.toSeq.sorted should equal(1004 to 1500) + } + + test("deserialization#complex without companion class") { + val json = + """{ + "entity_ids" : [ 1004, 1005, 1006, 1007, 1008, 1009, 1010, 1011, 1012, 1013, 1014, 1015, 1016, 1017, 1018, 1019, 1020, 1021, 1022, 1023, 1024, 1025, 1026, 1027, 1028, 1029, 1030, 1031, 1032, 1033, 1034, 1035, 1036, 1037, 1038, 1039, 1040, 1041, 1042, 1043, 1044, 1045, 1046, 1047, 1048, 1049, 1050, 1051, 1052, 1053, 1054, 1055, 1056, 1057, 1058, 1059, 1060, 1061, 1062, 1063, 1064, 1065, 1066, 1067, 1068, 1069, 1070, 1071, 1072, 1073, 1074, 1075, 1076, 1077, 1078, 1079, 1080, 1081, 1082, 1083, 1084, 1085, 1086, 1087, 1088, 1089, 1090, 1091, 1092, 1093, 1094, 1095, 1096, 1097, 1098, 1099, 1100, 1101, 1102, 1103, 1104, 1105, 1106, 1107, 1108, 1109, 1110, 1111, 1112, 1113, 1114, 1115, 1116, 1117, 1118, 1119, 1120, 1121, 1122, 1123, 1124, 1125, 1126, 1127, 1128, 1129, 1130, 1131, 1132, 1133, 1134, 1135, 1136, 1137, 1138, 1139, 1140, 1141, 1142, 1143, 1144, 1145, 1146, 1147, 1148, 1149, 1150, 1151, 1152, 1153, 1154, 1155, 1156, 1157, 1158, 1159, 1160, 1161, 1162, 1163, 1164, 1165, 1166, 1167, 1168, 1169, 1170, 1171, 1172, 1173, 1174, 1175, 1176, 1177, 1178, 1179, 1180, 1181, 1182, 1183, 1184, 1185, 1186, 1187, 1188, 1189, 1190, 1191, 1192, 1193, 1194, 1195, 1196, 1197, 1198, 1199, 1200, 1201, 1202, 1203, 1204, 1205, 1206, 1207, 1208, 1209, 1210, 1211, 1212, 1213, 1214, 1215, 1216, 1217, 1218, 1219, 1220, 1221, 1222, 1223, 1224, 1225, 1226, 1227, 1228, 1229, 1230, 1231, 1232, 1233, 1234, 1235, 1236, 1237, 1238, 1239, 1240, 1241, 1242, 1243, 1244, 1245, 1246, 1247, 1248, 1249, 1250, 1251, 1252, 1253, 1254, 1255, 1256, 1257, 1258, 1259, 1260, 1261, 1262, 1263, 1264, 1265, 1266, 1267, 1268, 1269, 1270, 1271, 1272, 1273, 1274, 1275, 1276, 1277, 1278, 1279, 1280, 1281, 1282, 1283, 1284, 1285, 1286, 1287, 1288, 1289, 1290, 1291, 1292, 1293, 1294, 1295, 1296, 1297, 1298, 1299, 1300, 1301, 1302, 1303, 1304, 1305, 1306, 1307, 1308, 1309, 1310, 1311, 1312, 1313, 1314, 1315, 1316, 1317, 1318, 1319, 1320, 1321, 1322, 1323, 1324, 1325, 1326, 1327, 1328, 1329, 1330, 1331, 1332, 1333, 1334, 1335, 1336, 1337, 1338, 1339, 1340, 1341, 1342, 1343, 1344, 1345, 1346, 1347, 1348, 1349, 1350, 1351, 1352, 1353, 1354, 1355, 1356, 1357, 1358, 1359, 1360, 1361, 1362, 1363, 1364, 1365, 1366, 1367, 1368, 1369, 1370, 1371, 1372, 1373, 1374, 1375, 1376, 1377, 1378, 1379, 1380, 1381, 1382, 1383, 1384, 1385, 1386, 1387, 1388, 1389, 1390, 1391, 1392, 1393, 1394, 1395, 1396, 1397, 1398, 1399, 1400, 1401, 1402, 1403, 1404, 1405, 1406, 1407, 1408, 1409, 1410, 1411, 1412, 1413, 1414, 1415, 1416, 1417, 1418, 1419, 1420, 1421, 1422, 1423, 1424, 1425, 1426, 1427, 1428, 1429, 1430, 1431, 1432, 1433, 1434, 1435, 1436, 1437, 1438, 1439, 1440, 1441, 1442, 1443, 1444, 1445, 1446, 1447, 1448, 1449, 1450, 1451, 1452, 1453, 1454, 1455, 1456, 1457, 1458, 1459, 1460, 1461, 1462, 1463, 1464, 1465, 1466, 1467, 1468, 1469, 1470, 1471, 1472, 1473, 1474, 1475, 1476, 1477, 1478, 1479, 1480, 1481, 1482, 1483, 1484, 1485, 1486, 1487, 1488, 1489, 1490, 1491, 1492, 1493, 1494, 1495, 1496, 1497, 1498, 1499, 1500 ], + "previous_cursor" : "$", + "next_cursor" : "2892e7ab37d44c6a15b438f78e8d76ed$" + }""" + val entityIdsResponse = parse[TestEntityIdsResponse](json) + entityIdsResponse.entityIds.sorted.size should be > 0 + } + + test("deserialization#complex with companion class") { + val json = + """{ + "entity_ids" : [ 1004, 1005, 1006, 1007, 1008, 1009, 1010, 1011, 1012, 1013, 1014, 1015, 1016, 1017, 1018, 1019, 1020, 1021, 1022, 1023, 1024, 1025, 1026, 1027, 1028, 1029, 1030, 1031, 1032, 1033, 1034, 1035, 1036, 1037, 1038, 1039, 1040, 1041, 1042, 1043, 1044, 1045, 1046, 1047, 1048, 1049, 1050, 1051, 1052, 1053, 1054, 1055, 1056, 1057, 1058, 1059, 1060, 1061, 1062, 1063, 1064, 1065, 1066, 1067, 1068, 1069, 1070, 1071, 1072, 1073, 1074, 1075, 1076, 1077, 1078, 1079, 1080, 1081, 1082, 1083, 1084, 1085, 1086, 1087, 1088, 1089, 1090, 1091, 1092, 1093, 1094, 1095, 1096, 1097, 1098, 1099, 1100, 1101, 1102, 1103, 1104, 1105, 1106, 1107, 1108, 1109, 1110, 1111, 1112, 1113, 1114, 1115, 1116, 1117, 1118, 1119, 1120, 1121, 1122, 1123, 1124, 1125, 1126, 1127, 1128, 1129, 1130, 1131, 1132, 1133, 1134, 1135, 1136, 1137, 1138, 1139, 1140, 1141, 1142, 1143, 1144, 1145, 1146, 1147, 1148, 1149, 1150, 1151, 1152, 1153, 1154, 1155, 1156, 1157, 1158, 1159, 1160, 1161, 1162, 1163, 1164, 1165, 1166, 1167, 1168, 1169, 1170, 1171, 1172, 1173, 1174, 1175, 1176, 1177, 1178, 1179, 1180, 1181, 1182, 1183, 1184, 1185, 1186, 1187, 1188, 1189, 1190, 1191, 1192, 1193, 1194, 1195, 1196, 1197, 1198, 1199, 1200, 1201, 1202, 1203, 1204, 1205, 1206, 1207, 1208, 1209, 1210, 1211, 1212, 1213, 1214, 1215, 1216, 1217, 1218, 1219, 1220, 1221, 1222, 1223, 1224, 1225, 1226, 1227, 1228, 1229, 1230, 1231, 1232, 1233, 1234, 1235, 1236, 1237, 1238, 1239, 1240, 1241, 1242, 1243, 1244, 1245, 1246, 1247, 1248, 1249, 1250, 1251, 1252, 1253, 1254, 1255, 1256, 1257, 1258, 1259, 1260, 1261, 1262, 1263, 1264, 1265, 1266, 1267, 1268, 1269, 1270, 1271, 1272, 1273, 1274, 1275, 1276, 1277, 1278, 1279, 1280, 1281, 1282, 1283, 1284, 1285, 1286, 1287, 1288, 1289, 1290, 1291, 1292, 1293, 1294, 1295, 1296, 1297, 1298, 1299, 1300, 1301, 1302, 1303, 1304, 1305, 1306, 1307, 1308, 1309, 1310, 1311, 1312, 1313, 1314, 1315, 1316, 1317, 1318, 1319, 1320, 1321, 1322, 1323, 1324, 1325, 1326, 1327, 1328, 1329, 1330, 1331, 1332, 1333, 1334, 1335, 1336, 1337, 1338, 1339, 1340, 1341, 1342, 1343, 1344, 1345, 1346, 1347, 1348, 1349, 1350, 1351, 1352, 1353, 1354, 1355, 1356, 1357, 1358, 1359, 1360, 1361, 1362, 1363, 1364, 1365, 1366, 1367, 1368, 1369, 1370, 1371, 1372, 1373, 1374, 1375, 1376, 1377, 1378, 1379, 1380, 1381, 1382, 1383, 1384, 1385, 1386, 1387, 1388, 1389, 1390, 1391, 1392, 1393, 1394, 1395, 1396, 1397, 1398, 1399, 1400, 1401, 1402, 1403, 1404, 1405, 1406, 1407, 1408, 1409, 1410, 1411, 1412, 1413, 1414, 1415, 1416, 1417, 1418, 1419, 1420, 1421, 1422, 1423, 1424, 1425, 1426, 1427, 1428, 1429, 1430, 1431, 1432, 1433, 1434, 1435, 1436, 1437, 1438, 1439, 1440, 1441, 1442, 1443, 1444, 1445, 1446, 1447, 1448, 1449, 1450, 1451, 1452, 1453, 1454, 1455, 1456, 1457, 1458, 1459, 1460, 1461, 1462, 1463, 1464, 1465, 1466, 1467, 1468, 1469, 1470, 1471, 1472, 1473, 1474, 1475, 1476, 1477, 1478, 1479, 1480, 1481, 1482, 1483, 1484, 1485, 1486, 1487, 1488, 1489, 1490, 1491, 1492, 1493, 1494, 1495, 1496, 1497, 1498, 1499, 1500 ], + "previous_cursor" : "$", + "next_cursor" : "2892e7ab37d44c6a15b438f78e8d76ed$" + }""" + val entityIdsResponse = parse[TestEntityIdsResponseWithCompanion](json) + entityIdsResponse.entityIds.sorted.size should be > 0 + } + + val json = """ + { + "map": { + "one": "two" + }, + "set": [1, 2, 3], + "string": "woo", + "list": [4, 5, 6], + "seq": [7, 8, 9], + "sequence": [10, 11, 12], + "collection": [13, 14, 15], + "indexed_seq": [16, 17, 18], + "random_access_seq": [19, 20, 21], + "vector": [22, 23, 24], + "big_decimal": 12.0, + "big_int": 13, + "int": 1, + "long": 2, + "char": "x", + "bool": false, + "short": 14, + "byte": 15, + "float": 34.5, + "double": 44.9, + "any": true, + "any_ref": "wah", + "int_map": { + "1": "1" + }, + "long_map": { + "2": 2 + } + }""" + + test( + "deserialization#case class with members of all ScalaSig types " + + "is parsable from a JSON object with those fields" + ) { + val got = parse[CaseClassWithAllTypes](json) + val expected = CaseClassWithAllTypes( + map = Map("one" -> "two"), + set = Set(1, 2, 3), + string = "woo", + list = List(4, 5, 6), + seq = Seq(7, 8, 9), + indexedSeq = IndexedSeq(16, 17, 18), + vector = Vector(22, 23, 24), + bigDecimal = BigDecimal("12.0"), + bigInt = 13, + int = 1, + long = 2L, + char = 'x', + bool = false, + short = 14, + byte = 15, + float = 34.5f, + double = 44.9d, + any = true, + anyRef = "wah", + intMap = Map(1 -> 1), + longMap = Map(2L -> 2L) + ) + + got should be(expected) + } + + test("deserialization# case class that throws an exception is not parsable from a JSON object") { + intercept[JsonMappingException] { + parse[CaseClassWithException]("""{}""") + } + } + + test("deserialization#case class nested inside of an object is parsable from a JSON object") { + parse[OuterObject.NestedCaseClass]("""{"id": 1}""") should be(OuterObject.NestedCaseClass(1)) + } + + test( + "deserialization#case class nested inside of an object nested inside of an object is parsable from a JSON object" + ) { + parse[OuterObject.InnerObject.SuperNestedCaseClass]("""{"id": 1}""") should be( + OuterObject.InnerObject.SuperNestedCaseClass(1) + ) + } + + test("A case class with array members is parsable from a JSON object") { + val jsonStr = """ + { + "one":"1", + "two":["a","b","c"], + "three":[1,2,3], + "four":[4, 5], + "five":["x", "y"], + "bools":["true", false], + "bytes":[1,2], + "doubles":[1,5.0], + "floats":[1.1, 22] + } + """ + + val c = parse[CaseClassWithArrays](jsonStr) + c.one should be("1") + c.two should be(Array("a", "b", "c")) + c.three should be(Array(1, 2, 3)) + c.four should be(Array(4L, 5L)) + c.five should be(Array('x', 'y')) + + JsonDiff.assertDiff( + """{"bools":[true,false],"bytes":"AQI=","doubles":[1.0,5.0],"five":"xy","floats":[1.1,22.0],"four":[4,5],"one":"1","three":[1,2,3],"two":["a","b","c"]}""", + generate(c) + ) + } + + test("deserialization#case class with collection of Longs array of longs") { + val c = parse[CaseClassWithArrayLong]("""{"array":[3,1,2]}""") + c.array.sorted should equal(Array(1, 2, 3)) + } + + test("deserialization#case class with collection of Longs seq of longs") { + val c = parse[CaseClassWithSeqLong]("""{"seq":[3,1,2]}""") + c.seq.sorted should equal(Seq(1, 2, 3)) + } + + test("deserialization#case class with an ArrayList of Integers") { + val c = parse[CaseClassWithArrayListOfIntegers]("""{"arraylist":[3,1,2]}""") + val l = new java.util.ArrayList[Integer](3) + l.add(3) + l.add(1) + l.add(2) + c.arraylist should equal(l) + } + + test("serde#case class with a SortedMap[String, Int]") { + val origCaseClass = CaseClassWithSortedMap(scala.collection.SortedMap("aggregate" -> 20)) + val caseClassJson = generate(origCaseClass) + val caseClass = parse[CaseClassWithSortedMap](caseClassJson) + caseClass should equal(origCaseClass) + } + + test("serde#case class with a Seq of Longs") { + val origCaseClass = CaseClassWithSeqOfLongs(Seq(10, 20, 30)) + val caseClassJson = generate(origCaseClass) + val caseClass = parse[CaseClassWithSeqOfLongs](caseClassJson) + caseClass should equal(origCaseClass) + } + + test("deserialization#seq of longs") { + val seq = parse[Seq[Long]]("""[3,1,2]""") + seq.sorted should equal(Seq(1L, 2L, 3L)) + } + + test("deserialization#parse seq of longs") { + val ids = parse[Seq[Long]]("[3,1,2]") + ids.sorted should equal(Seq(1L, 2L, 3L)) + } + + test("deserialization#handle options and defaults in case class") { + val bob = parse[Person](""" + { + "id" :1, + "name" : "Bob", + "age" : 21 + } + """) + bob should equal(Person(1, "Bob", None, Some(21), None)) + } + + test("deserialization#missing required field") { + intercept[CaseClassMappingException] { + parse[Person](""" + { + } + """) + } + } + + test("deserialization#incorrectly specified required field") { + intercept[CaseClassMappingException] { + parse[PersonWithThings](""" + { + "id" :1, + "name" : "Bob", + "age" : 21, + "things" : { + "foo" : [ + "IhaveNoKey" + ] + } + } + """) + } + } + + test("serialization#nulls will not render") { + generate(Person(1, null, null, null)) should equal("""{"id":1,"nickname":"unknown"}""") + } + + test("deserialization#string wrapper deserialization") { + val parsedValue = parse[ObjWithTestId](""" + { + "id": "5" + } + """) + val expectedValue = ObjWithTestId(TestIdStringWrapper("5")) + + parsedValue should equal(expectedValue) + parsedValue.id.onlyValue should equal(expectedValue.id.onlyValue) + parsedValue.id.asString should equal(expectedValue.id.asString) + parsedValue.id.toString should equal(expectedValue.id.toString) + } + + test("deserialization#parse input stream") { + val is = new ByteArrayInputStream("""{"foo": "bar"}""".getBytes) + mapper.parse[Blah](is) should equal(Blah("bar")) + } + + test("serialization#ogging Trait fields should be ignored") { + generate(Group3("123")) should be("""{"id":"123"}""") + } + + test("deserialization#class with no constructor") { + parse[NoConstructorArgs]("""{}""") + } + + //Jackson parses numbers into boolean type without error. see https://jira.codehaus.org/browse/JACKSON-78 + test("deserialization#case class with boolean as number") { + parse[CaseClassWithBoolean](""" { + "foo": 100 + }""") should equal(CaseClassWithBoolean(true)) + } + + //Jackson parses numbers into boolean type without error. see https://jira.codehaus.org/browse/JACKSON-78 + test("deserialization#case class with Seq[Boolean]") { + parse[CaseClassWithSeqBooleans](""" { + "foos": [100, 5, 0, 9] + }""") should equal(CaseClassWithSeqBooleans(Seq(true, true, false, true))) + } + + //Jackson parses numbers into boolean type without error. see https://jira.codehaus.org/browse/JACKSON-78 + test("Seq[Boolean]") { + parse[Seq[Boolean]]("""[100, 5, 0, 9]""") should equal(Seq(true, true, false, true)) + } + + test("case class with boolean as number 0") { + parse[CaseClassWithBoolean](""" { + "foo": 0 + }""") should equal(CaseClassWithBoolean(false)) + } + + test("case class with boolean as string") { + assertJsonParse[CaseClassWithBoolean]( + """ { + "foo": "bar" + }""", + withErrors = Seq("foo: 'bar' is not a valid Boolean")) + } + + test("case class with boolean number as string") { + assertJsonParse[CaseClassWithBoolean]( + """ { + "foo": "1" + }""", + withErrors = Seq("foo: '1' is not a valid Boolean")) + } + + val msgHiJsonStr = """{"msg":"hi"}""" + + test("parse jsonParser") { + val jsonNode = mapper.parse[JsonNode]("{}") + val jsonParser = new TreeTraversingParser(jsonNode) + mapper.parse[JsonNode](jsonParser) should equal(jsonNode) + } + + test("writeValue") { + val os = new ByteArrayOutputStream() + mapper.writeValue(Map("msg" -> "hi"), os) + os.close() + new String(os.toByteArray) should equal(msgHiJsonStr) + } + + test("writeValueAsBuf") { + val buf = mapper.writeValueAsBuf(Map("msg" -> "hi")) + val Buf.Utf8(str) = buf + str should equal(msgHiJsonStr) + } + + test("writeStringMapAsBuf") { + val buf = mapper.writeStringMapAsBuf(Map("msg" -> "hi")) + val Buf.Utf8(str) = buf + str should equal(msgHiJsonStr) + } + + test("writePrettyString") { + val jsonStr = mapper.writePrettyString("""{"msg": "hi"}""") + mapper.parse[JsonNode](jsonStr).get("msg").textValue() should equal("hi") + } + + test("reader") { + assert(mapper.reader[JsonNode] != null) + } + + test( + "deserialization#jackson JsonDeserialize annotations deserializes json to case class with 2 decimal places for mandatory field" + ) { + parse[CaseClassWithCustomDecimalFormat](""" { + "my_big_decimal": 23.1201 + }""") should equal(CaseClassWithCustomDecimalFormat(BigDecimal(23.12), None)) + } + + test("deserialization#jackson JsonDeserialize annotations long with JsonDeserialize") { + val result = parse[CaseClassWithLongAndDeserializer](""" { + "long": 12345 + }""") + result should equal(CaseClassWithLongAndDeserializer(12345)) + result.long should equal(12345L) // promotes the returned integer into a long + } + + test( + "deserialization#jackson JsonDeserialize annotations deserializes json to case class with 2 decimal places for option field" + ) { + parse[CaseClassWithCustomDecimalFormat](""" { + "my_big_decimal": 23.1201, + "opt_my_big_decimal": 23.1201 + }""") should equal( + CaseClassWithCustomDecimalFormat(BigDecimal(23.12), Some(BigDecimal(23.12))) + ) + } + + test("deserialization#jackson JsonDeserialize annotations opt long with JsonDeserialize") { + parse[CaseClassWithOptionLongAndDeserializer](""" + { + "opt_long": 12345 + } + """) should equal(CaseClassWithOptionLongAndDeserializer(Some(12345))) + } + + test("deserialization#case class in package object") { + parse[SimplePersonInPackageObject]("""{"name": "Steve"}""") should equal( + SimplePersonInPackageObject("Steve") + ) + } + + test("deserialization#case class in package object uses default when name not specified") { + parse[SimplePersonInPackageObject]("""{}""") should equal( + SimplePersonInPackageObject() + ) + } + + test( + "deserialization#case class in package object without constructor params and parsing an empty json object") { + parse[SimplePersonInPackageObjectWithoutConstructorParams]("""{}""") should equal( + SimplePersonInPackageObjectWithoutConstructorParams() + ) + } + + test( + "deserialization#case class in package object without constructor params and parsing a json object with extra fields" + ) { + parse[SimplePersonInPackageObjectWithoutConstructorParams]( + """{"name": "Steve"}""") should equal( + SimplePersonInPackageObjectWithoutConstructorParams() + ) + } + + test("serialization#upport sealed traits and case objects#json serialization") { + val vin = Random.alphanumeric.take(17).mkString + val vehicle = Vehicle(vin, Audi) + + mapper.writeValueAsString(vehicle) should equal(s"""{"vin":"$vin","type":"audi"}""") + } + + test("deserialization#json creator") { + parse[TestJsonCreator]("""{"s" : "1234"}""") should equal(TestJsonCreator("1234")) + } + + test("deserialization#json creator with parameterized type") { + parse[TestJsonCreator2]("""{"strings" : ["1", "2", "3", "4"]}""") should equal( + TestJsonCreator2(Seq(1, 2, 3, 4))) + } + + test("deserialization#multiple case class constructors#no annotation") { + intercept[AssertionError] { + parse[CaseClassWithMultipleConstructors]( + """ + |{ + | "number1" : 12345, + | "number2" : 65789, + | "number3" : 99999 + |} + |""".stripMargin + ) + } + } + + test("deserialization#multiple case class constructors#annotated") { + parse[CaseClassWithMultipleConstructorsAnnotated]( + """ + |{ + | "number_as_string1" : "12345", + | "number_as_string2" : "65789", + | "number_as_string3" : "99999" + |} + |""".stripMargin + ) should equal( + CaseClassWithMultipleConstructorsAnnotated(12345L, 65789L, 99999L) + ) + } + + test("deserialization#JsonFormat with c.t.util.Time") { + val timeString = "2018-09-14T23:20:08.000-07:00" + val format = new com.twitter.util.TimeFormat( + "yyyy-MM-dd'T'HH:mm:ss.SSSXXX", + None, + TimeZone.getTimeZone("UTC")) + val time: com.twitter.util.Time = format.parse(timeString) + + parse[TimeWithFormat]( + s""" + |{ + | "when": "$timeString" + |} + |""".stripMargin + ) should equal(TimeWithFormat(time)) + } + + test("deserialization#complex object") { + val profiles = Seq( + LimiterProfile(1, "up"), + LimiterProfile(2, "down"), + LimiterProfile(3, "left"), + LimiterProfile(4, "right")) + val expected = LimiterProfiles(profiles) + + parse[LimiterProfiles](""" + |{ + | "profiles" : [ + | { + | "id" : "1", + | "name" : "up" + | }, + | { + | "id" : "2", + | "name" : "down" + | }, + | { + | "id" : "3", + | "name" : "left" + | }, + | { + | "id" : "4", + | "name" : "right" + | } + | ] + |} + |""".stripMargin) should equal(expected) + } + + test("deserialization#simple class with boxed primitive constructor args") { + parse[CaseClassWithBoxedPrimitives](""" + |{ + | "events" : 2, + | "errors" : 0 + |} + |""".stripMargin) should equal(CaseClassWithBoxedPrimitives(2, 0)) + } + + test("deserialization#complex object with primitives") { + val json = """ + |{ + | "cluster_name" : "cluster1", + | "zone" : "abcd1", + | "environment" : "devel", + | "job" : "devel-1", + | "owners" : "owning-group", + | "dtab" : "/s#/devel/abcd1/space/devel-1", + | "address" : "local-test.foo.bar.com:1111/owning-group/space/devel-1", + | "enabled" : "true" + |} + |""".stripMargin + + parse[AddClusterRequest](json) should equal( + AddClusterRequest( + clusterName = "cluster1", + zone = "abcd1", + environment = "devel", + job = "devel-1", + dtab = "/s#/devel/abcd1/space/devel-1", + address = "local-test.foo.bar.com:1111/owning-group/space/devel-1", + owners = "owning-group" + )) + } + + test("deserialization#inject request field fails with mapping exception") { + // with an injector and default injectable values but the field cannot be parsed + // thus we get a case class mapping exception + // without an injector, the incoming JSON is used to populate the case class, which + // is empty and thus we also get a case class mapping exception + intercept[CaseClassMappingException] { + parse[ClassWithQueryParamDateTimeInject]("""{}""") + } + } + + test("deserialization#ignored constructor field - no default value") { + val json = + """ + |{ + | "name" : "Widget", + | "description" : "This is a thing." + |} + |""".stripMargin + + val e = intercept[CaseClassMappingException] { + parse[CaseClassIgnoredFieldInConstructorNoDefault](json) + } + e.errors.head.getMessage should be("id: ignored field has no default value specified") + } + + test("ignored constructor field - with default value") { + val json = + """ + |{ + | "name" : "Widget", + | "description" : "This is a thing." + |} + |""".stripMargin + + val result = parse[CaseClassIgnoredFieldInConstructorWithDefault](json) + result.id should equal(42L) + result.name should equal("Widget") + result.description should equal("This is a thing.") + } + + test("deserialization#mixin annotations") { + val points = Points(first = Point(1, 1), second = Point(4, 5)) + val json = + """ + |{ + | "first": { "x": 1, "y": 1 }, + | "second": { "x": 4, "y": 5 } + |} + |""".stripMargin + + parse[Points](json) should equal(points) + generate(points) should be("""{"first":{"x":1,"y":1},"second":{"x":4,"y":5}}""".stripMargin) + } + + test("deserialization#ignore type with no default fails") { + val json = + """ + |{ + | "name" : "Widget", + | "description" : "This is a thing." + |} + |""".stripMargin + val e = intercept[CaseClassMappingException] { + parse[ContainsAnIgnoreTypeNoDefault](json) + } + e.errors.head.getMessage should be("ignored: ignored field has no default value specified") + } + + test("deserialization#ignore type with default passes") { + val json = + """ + |{ + | "name" : "Widget", + | "description" : "This is a thing." + |} + |""".stripMargin + val result = parse[ContainsAnIgnoreTypeWithDefault](json) + result.ignored should equal(IgnoreMe(42L)) + result.name should equal("Widget") + result.description should equal("This is a thing.") + } + + test("deserialization#parse SimpleClassWithInjection fails") { + // no field injection configured, parsing happens normally + val json = + """{ + | "hello" : "Mom" + |}""".stripMargin + + val result = parse[SimpleClassWithInjection](json) + result.hello should equal("Mom") + } + + test("deserialization#support JsonView") { + /* using [[Item]] which has @JsonView annotation on all fields: + Public -> id + Public -> name + Internal -> owner + */ + val json = + """ + |{ + | "id": 42, + | "name": "My Item" + |} + |""".stripMargin + + val item = + mapper.underlying.readerWithView[Views.Public].forType(classOf[Item]).readValue[Item](json) + item.id should equal(42L) + item.name should equal("My Item") + item.owner should equal("") + } + + test("deserialization#support JsonView 2") { + /* using [[Item]] which has @JsonView annotation on all fields: + Public -> id + Public -> name + Internal -> owner + */ + val json = + """ + |{ + | "owner": "Mister Magoo" + |} + |""".stripMargin + + val item = + mapper.underlying.readerWithView[Views.Internal].forType(classOf[Item]).readValue[Item](json) + item.id should equal(1L) + item.name should equal("") + item.owner should equal("Mister Magoo") + } + + test("deserialization#support JsonView 3") { + /* using [[ItemSomeViews]] which has no @JsonView annotation on owner field: + Public -> id + Public -> name + -> owner + */ + val json = + """ + |{ + | "id": 42, + | "name": "My Item", + | "owner": "Mister Magoo" + |} + |""".stripMargin + + val item = + mapper.underlying + .readerWithView[Views.Public] + .forType(classOf[ItemSomeViews]) + .readValue[ItemSomeViews](json) + item.id should equal(42L) + item.name should equal("My Item") + item.owner should equal("Mister Magoo") + } + + test("deserialization#support JsonView 4") { + /* using [[ItemSomeViews]] which has no @JsonView annotation on owner field: + Public -> id + Public -> name + -> owner + */ + val json = + """ + |{ + | "owner": "Mister Magoo" + |} + |""".stripMargin + + val item = + mapper.underlying + .readerWithView[Views.Internal] + .forType(classOf[ItemSomeViews]) + .readValue[ItemSomeViews](json) + item.id should equal(1L) + item.name should equal("") + item.owner should equal("Mister Magoo") + } + + test("deserialization#support JsonView 5") { + /* using [[ItemNoDefaultForView]] which has @JsonView annotation on non-defaulted field: + Public -> name (no default value) + */ + val json = + """ + |{ + | "name": "My Item" + |} + |""".stripMargin + + val item = + mapper.underlying + .readerWithView[Views.Public] + .forType(classOf[ItemNoDefaultForView]) + .readValue[ItemNoDefaultForView](json) + item.name should equal("My Item") + } + + test("deserialization#support JsonView 6") { + /* using [[ItemNoDefaultForView]] which has @JsonView annotation on non-defaulted field: + Public -> name (no default value) + */ + + intercept[CaseClassMappingException] { + mapper.underlying + .readerWithView[Views.Public] + .forType(classOf[ItemNoDefaultForView]) + .readValue[ItemNoDefaultForView]("{}") + } + } + + test("deserialization#test @JsonNaming") { + val json = + """ + |{ + | "please-use-kebab-case": true + |} + |""".stripMargin + + parse[CaseClassWithKebabCase](json) should equal(CaseClassWithKebabCase(true)) + + val snakeCaseMapper = ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying) + snakeCaseMapper.parse[CaseClassWithKebabCase](json) should equal(CaseClassWithKebabCase(true)) + + val camelCaseMapper = ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying) + camelCaseMapper.parse[CaseClassWithKebabCase](json) should equal(CaseClassWithKebabCase(true)) + } + + test("deserialization#test @JsonNaming 2") { + val snakeCaseMapper = ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying) + + // case class is marked with @JsonNaming and the default naming strategy is LOWER_CAMEL_CASE + // which overrides the mapper's configured naming strategy + val json = + """ + |{ + | "thisFieldShouldUseDefaultPropertyNamingStrategy": true + |} + |""".stripMargin + + snakeCaseMapper.parse[UseDefaultNamingStrategy](json) should equal( + UseDefaultNamingStrategy(true)) + + // default for mapper is snake_case, but the case class is annotated with @JsonNaming which + // tells the mapper what property naming strategy to expect. + parse[UseDefaultNamingStrategy](json) should equal(UseDefaultNamingStrategy(true)) + } + + test("deserialization#test @JsonNaming 3") { + val camelCaseMapper = ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying) + + val json = + """ + |{ + | "thisFieldShouldUseDefaultPropertyNamingStrategy": true + |} + |""".stripMargin + + camelCaseMapper.parse[UseDefaultNamingStrategy](json) should equal( + UseDefaultNamingStrategy(true)) + + // default for mapper is snake_case, but the case class is annotated with @JsonNaming which + // tells the mapper what property naming strategy to expect. + parse[UseDefaultNamingStrategy](json) should equal(UseDefaultNamingStrategy(true)) + } + + test("deserialization#test @JsonNaming with mixin") { + val json = + """ + |{ + | "will-this-get-the-right-casing": true + |} + |""".stripMargin + + parse[CaseClassShouldUseKebabCaseFromMixin](json) should equal( + CaseClassShouldUseKebabCaseFromMixin(true)) + + val snakeCaseMapper = ScalaObjectMapper.snakeCaseObjectMapper(mapper.underlying) + snakeCaseMapper.parse[CaseClassShouldUseKebabCaseFromMixin](json) should equal( + CaseClassShouldUseKebabCaseFromMixin(true)) + + val camelCaseMapper = ScalaObjectMapper.camelCaseObjectMapper(mapper.underlying) + camelCaseMapper.parse[CaseClassShouldUseKebabCaseFromMixin](json) should equal( + CaseClassShouldUseKebabCaseFromMixin(true)) + } + + test("serde#can serialize and deserialize polymorphic types") { + val view = View(shapes = Seq(Circle(10), Rectangle(5, 5), Circle(3), Rectangle(10, 5))) + val result = generate(view) + result should be( + """{"shapes":[{"type":"circle","radius":10},{"type":"rectangle","width":5,"height":5},{"type":"circle","radius":3},{"type":"rectangle","width":10,"height":5}]}""") + + parse[View](result) should equal(view) + } + + test("serde#can serialize and deserialize optional polymorphic types") { + val view = + OptionalView(shapes = Seq(Circle(10), Rectangle(5, 5)), optional = Some(Rectangle(10, 5))) + val result = generate(view) + result should be( + """{"shapes":[{"type":"circle","radius":10},{"type":"rectangle","width":5,"height":5}],"optional":{"type":"rectangle","width":10,"height":5}}""") + + parse[OptionalView](result) should equal(view) + } + + test("serde#can deserialize case class with option when missing") { + val view = OptionalView(shapes = Seq(Circle(10), Rectangle(5, 5)), optional = None) + val result = generate(view) + result should equal( + """{"shapes":[{"type":"circle","radius":10},{"type":"rectangle","width":5,"height":5}]}""") + + parse[OptionalView](result) should equal(view) + } + + test("nullable deserializer - allows parsing a JSON null into a null type") { + parse[WithNullableCarMake]( + """{"non_null_car_make": "vw", "nullable_car_make": null}""") should be( + WithNullableCarMake(nonNullCarMake = CarMakeEnum.vw, nullableCarMake = null)) + } + + test("nullable deserializer - fail on missing value even when null allowed") { + val e = intercept[CaseClassMappingException] { + parse[WithNullableCarMake]("""{}""") + } + e.errors.size should equal(2) + e.errors.head.getMessage should equal("non_null_car_make: field is required") + e.errors.last.getMessage should equal("nullable_car_make: field is required") + } + + test("nullable deserializer - allow null values when default is provided") { + parse[WithDefaultNullableCarMake]("""{"nullable_car_make":null}""") should be( + WithDefaultNullableCarMake(null)) + } + + test("nullable deserializer - use default when null is allowed, but not provided") { + parse[WithDefaultNullableCarMake]("""{}""") should be( + WithDefaultNullableCarMake(CarMakeEnum.ford)) + } + + test("deserialization of json array into a string field") { + val json = + """ + |{ + | "value": [{ + | "languages": "es", + | "reviewable": "tweet", + | "baz": "too", + | "foo": "bar", + | "coordinates": "polar" + | }] + |} + |""".stripMargin + + val expected = WithJsonStringType( + """[{"languages":"es","reviewable":"tweet","baz":"too","foo":"bar","coordinates":"polar"}]""".stripMargin) + parse[WithJsonStringType](json) should equal(expected) + } + + test("deserialization of json int into a string field") { + val json = + """ + |{ + | "value": 1 + |} + |""".stripMargin + + val expected = WithJsonStringType("1") + parse[WithJsonStringType](json) should equal(expected) + } + + test("deserialization of json object into a string field") { + val json = + """ + |{ + | "value": {"foo": "bar"} + |} + |""".stripMargin + + val expected = WithJsonStringType("""{"foo":"bar"}""") + parse[WithJsonStringType](json) should equal(expected) + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraint.java b/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraint.java new file mode 100644 index 0000000000..0e8b538e22 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraint.java @@ -0,0 +1,38 @@ +package com.twitter.util.jackson; + +import java.lang.annotation.Documented; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {ThrowsRuntimeExceptionConstraintValidator.class}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +public @interface ThrowsRuntimeExceptionConstraint { + + /** Every constraint annotation must define a message element of type String. */ + String message() default ""; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. + */ + Class[] payload() default {}; +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraintValidator.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraintValidator.scala new file mode 100644 index 0000000000..741697bcaa --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/ThrowsRuntimeExceptionConstraintValidator.scala @@ -0,0 +1,14 @@ +package com.twitter.util.jackson + +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import scala.util.control.NoStackTrace + +class ThrowsRuntimeExceptionConstraintValidator + extends ConstraintValidator[ThrowsRuntimeExceptionConstraint, Any] { + override def isValid( + obj: Any, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + throw new RuntimeException("validator foo error") with NoStackTrace + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/ValidationScalaObjectMapperTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/ValidationScalaObjectMapperTest.scala new file mode 100644 index 0000000000..18eb547cea --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/ValidationScalaObjectMapperTest.scala @@ -0,0 +1,244 @@ +package com.twitter.util.jackson + +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.ValidationError +import com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException +import java.util.UUID +import org.junit.runner.RunWith +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class ValidationScalaObjectMapperTest + extends AnyFunSuite + with BeforeAndAfterAll + with Matchers + with ScalaObjectMapperFunctions { + + /* Class under test */ + protected val mapper: ScalaObjectMapper = ScalaObjectMapper() + + override def beforeAll(): Unit = { + super.beforeAll() + mapper.registerModule(ScalaObjectMapperTest.MixInAnnotationsModule) + } + + test("fail when CaseClassWithSeqWrappedValueLongWithValidation with null array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqWrappedValueLongWithValidation]("""{"seq": [null]}""") + } + } + + test("fail when CaseClassWithSeqWrappedValueLongWithValidation with invalid field") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqWrappedValueLongWithValidation]("""{"seq": [{"value": 0}]}""") + } + } + + test("fail when CaseClassWithSeqOfCaseClassWithValidation with null array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqOfCaseClassWithValidation]("""{"seq": [null]}""") + } + } + + test("fail when CaseClassWithSeqOfCaseClassWithValidation with invalid array element") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqOfCaseClassWithValidation]("""{"seq": [0]}""") + } + } + + test("fail when CaseClassWithSeqOfCaseClassWithValidation with null field in object") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqOfCaseClassWithValidation]("""{"seq": [{"value": null}]}""") + } + } + + test("fail when CaseClassWithSeqOfCaseClassWithValidation with invalid field in object") { + intercept[CaseClassMappingException] { + parse[CaseClassWithSeqOfCaseClassWithValidation]("""{"seq": [{"value": 0}]}""") + } + } + + test( + "A case class with an Option[Int] member and validation annotation is parsable from a JSON object with value that passes the validation") { + parse[CaseClassWithOptionAndValidation]("""{ "towing_capacity": 10000 }""") should be( + CaseClassWithOptionAndValidation(Some(10000))) + } + + test( + "A case class with an Option[Int] member and validation annotation is parsable from a JSON object with null value") { + parse[CaseClassWithOptionAndValidation]("""{ "towing_capacity": null }""") should be( + CaseClassWithOptionAndValidation(None)) + } + + test( + "A case class with an Option[Int] member and validation annotation is parsable from a JSON without that field") { + parse[CaseClassWithOptionAndValidation]("""{}""") should be( + CaseClassWithOptionAndValidation(None)) + } + + test( + "A case class with an Option[Int] member and validation annotation is parsable from a JSON object with value that fails the validation") { + val e = intercept[CaseClassMappingException] { + parse[CaseClassWithOptionAndValidation]("""{ "towing_capacity": 1 }""") should be( + CaseClassWithOptionAndValidation(Some(1))) + } + e.errors.head.getMessage should be("towing_capacity: must be greater than or equal to 100") + } + + test( + "A case class with an Option[Boolean] member and incompatible validation annotation is not parsable from a JSON object") { + val e = intercept[jakarta.validation.UnexpectedTypeException] { + parse[CaseClassWithOptionAndIncompatibleValidation]("""{ "are_you": true }""") should be( + CaseClassWithOptionAndIncompatibleValidation(Some(true))) + } + } + + test( + "A case class with an Option[Boolean] member and incompatible validation annotation is parsable from a JSON object with null value") { + parse[CaseClassWithOptionAndIncompatibleValidation]("""{ "are_you": null }""") should be( + CaseClassWithOptionAndIncompatibleValidation(None)) + } + + test( + "A case class with an Option[Boolean] member and incompatible validation annotation is parsable from a JSON object without that field") { + parse[CaseClassWithOptionAndIncompatibleValidation]("""{}""") should be( + CaseClassWithOptionAndIncompatibleValidation(None)) + } + + test("deserialization#JsonCreator with Validation") { + val e = intercept[CaseClassMappingException] { + // fails validation + parse[TestJsonCreatorWithValidation]("""{ "s": "" }""") + } + e.errors.size should equal(1) + e.errors.head.getMessage should be("s: must not be empty") + } + + test("deserialization#JsonCreator with Validation 1") { + // works with multiple validations + var e = intercept[CaseClassMappingException] { + // fails validation + parse[TestJsonCreatorWithValidations]("""{ "s": "" }""") + } + e.errors.size should equal(2) + // errors are alpha-sorted by message to be stable for testing + e.errors.head.getMessage should be("s: not one of [42, 137]") + e.errors.last.getMessage should be("s: must not be empty") + + e = intercept[CaseClassMappingException] { + // fails validation + parse[TestJsonCreatorWithValidations]("""{ "s": "99" }""") + } + e.errors.size should equal(1) + e.errors.head.getMessage should be("s: 99 not one of [42, 137]") + + parse[TestJsonCreatorWithValidations]("""{ "s": "42" }""") should equal( + TestJsonCreatorWithValidations("42")) + + parse[TestJsonCreatorWithValidations]("""{ "s": "137" }""") should equal( + TestJsonCreatorWithValidations(137)) + } + + test("deserialization#JsonCreator with Validation 2") { + val uuid = UUID.randomUUID().toString + + var e = intercept[CaseClassMappingException] { + // fails validation + parse[CaseClassWithMultipleConstructorsAnnotatedAndValidations]( + s""" + |{ + | "number_as_string1" : "", + | "number_as_string2" : "20002", + | "third_argument" : "$uuid" + |} + |""".stripMargin + ) + } + e.errors.size should equal(1) + e.errors.head.getMessage should be("number_as_string1: must not be empty") + + e = intercept[CaseClassMappingException] { + // fails validation + parse[CaseClassWithMultipleConstructorsAnnotatedAndValidations]( + s""" + |{ + | "number_as_string1" : "", + | "number_as_string2" : "65789", + | "third_argument" : "$uuid" + |} + |""".stripMargin + ) + } + e.errors.size should equal(2) + // errors are alpha-sorted by message to be stable for testing + e.errors.head.getMessage should be("number_as_string1: must not be empty") + e.errors.last.getMessage should be("number_as_string2: 65789 not one of [10001, 20002, 30003]") + + e = intercept[CaseClassMappingException] { + // fails validation + parse[CaseClassWithMultipleConstructorsAnnotatedAndValidations]( + s""" + |{ + | "number_as_string1" : "", + | "number_as_string2" : "65789", + | "third_argument" : "foobar" + |} + |""".stripMargin + ) + } + e.errors.size should equal(3) + // errors are alpha-sorted by message to be stable for testing + e.errors.head.getMessage should be("number_as_string1: must not be empty") + e.errors(1).getMessage should be("number_as_string2: 65789 not one of [10001, 20002, 30003]") + e.errors(2).getMessage should be("third_argument: must be a valid UUID") + + parse[CaseClassWithMultipleConstructorsAnnotatedAndValidations]( + s""" + |{ + | "number_as_string1" : "12345", + | "number_as_string2" : "20002", + | "third_argument" : "$uuid" + |} + |""".stripMargin + ) should equal( + CaseClassWithMultipleConstructorsAnnotatedAndValidations(12345L, 20002L, uuid) + ) + } + + test("deserialization#mixin annotations with validations") { + val json = + """ + |{ + | "first": { "x": -1, "y": 120 }, + | "second": { "x": 4, "y": 5 } + |} + |""".stripMargin + + val e = intercept[CaseClassMappingException] { + parse[Points](json) + } + e.errors.size should equal(2) + e.errors.head.getMessage should be("first.x: must be greater than or equal to 0") + e.errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Field, None) => + violation.getPropertyPath.toString should equal("Point.x") + violation.getMessage should equal("must be greater than or equal to 0") + violation.getInvalidValue should equal(-1) + violation.getRootBeanClass should equal(classOf[Point]) + violation.getRootBean == null should be(true) + case _ => fail() + } + e.errors.last.getMessage should be("first.y: must be less than or equal to 100") + e.errors.last.reason.detail match { + case ValidationError(violation, ValidationError.Field, None) => + violation.getPropertyPath.toString should equal("Point.y") + violation.getMessage should equal("must be less than or equal to 100") + violation.getInvalidValue should equal(120) + violation.getRootBeanClass should equal(classOf[Point]) + violation.getRootBean == null should be(true) + case _ => fail() + } + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/YAMLTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/YAMLTest.scala new file mode 100644 index 0000000000..79d9b16f25 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/YAMLTest.scala @@ -0,0 +1,159 @@ +package com.twitter.util.jackson + +import com.twitter.io.Buf +import java.io.{ByteArrayInputStream, File => JFile} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class YAMLTest extends AnyFunSuite with Matchers with FileResources { + + test("YAML#parse 1") { + YAML.parse[Map[String, Int]](""" + |--- + |a: 1 + |b: 2 + |c: 3""".stripMargin) match { + case Some(map) => + map should be(Map("a" -> 1, "b" -> 2, "c" -> 3)) + case _ => fail() + } + } + + test("YAML#parse 2") { + // without quotes + YAML.parse[Seq[String]](""" + |--- + |- a + |- b + |- c""".stripMargin) match { + case Some(Seq("a", "b", "c")) => // pass + case _ => fail() + } + // with quotes + YAML.parse[Seq[String]](""" + |--- + |- "a" + |- "b" + |- "c"""".stripMargin) match { + case Some(Seq("a", "b", "c")) => // pass + case _ => fail() + } + } + + test("YAML#parse 3") { + // without quotes + YAML.parse[FooClass](""" + |--- + |id: abcde1234""".stripMargin) match { + case Some(foo) => + foo.id should equal("abcde1234") + case _ => fail() + } + // with quotes + YAML.parse[FooClass](""" + |--- + |id: "abcde1234"""".stripMargin) match { + case Some(foo) => + foo.id should equal("abcde1234") + case _ => fail() + } + } + + test("YAML#parse 4") { + YAML.parse[CaseClassWithNotEmptyValidation](""" + |--- + |name: '' + |make: vw""".stripMargin) match { + case Some(clazz) => + // performs no validations + clazz.name.isEmpty should be(true) + clazz.make should equal(CarMakeEnum.vw) + case _ => fail() + } + } + + test("YAML#parse 5") { + val buf = Buf.Utf8("""{"id": "abcd1234"}""") + YAML.parse[FooClass](buf) match { + case Some(foo) => + foo.id should equal("abcd1234") + case _ => fail() + } + } + + test("YAML#parse 6") { + val inputStream = new ByteArrayInputStream("""{"id": "abcd1234"}""".getBytes("UTF-8")) + try { + YAML.parse[FooClass](inputStream) match { + case Some(foo) => + foo.id should equal("abcd1234") + case _ => fail() + } + } finally { + inputStream.close() + } + } + + test("YAML#parse 7") { + withTempFolder { + // without quotes + val file1: JFile = + writeStringToFile( + folderName, + "test-file-quotes", + ".yaml", + """|--- + |id: 999999999""".stripMargin) + YAML.parse[FooClass](file1) match { + case Some(foo) => + foo.id should equal("999999999") + case _ => fail() + } + // with quotes + val file2: JFile = + writeStringToFile( + folderName, + "test-file-noquotes", + ".yaml", + """|--- + |id: "999999999"""".stripMargin) + YAML.parse[FooClass](file2) match { + case Some(foo) => + foo.id should equal("999999999") + case _ => fail() + } + } + } + + test("YAML#write 1") { + YAML.write(Map("a" -> 1, "b" -> 2)) should equal("""|--- + |a: 1 + |b: 2 + |""".stripMargin) + } + + test("YAML#write 2") { + YAML.write(Seq("a", "b", "c")) should equal("""|--- + |- "a" + |- "b" + |- "c" + |""".stripMargin) + } + + test("YAML#write 3") { + YAML.write(FooClass("abcd1234")) should equal("""|--- + |id: "abcd1234" + |""".stripMargin) + } + + test("YAML.Resource#parse resource") { + YAML.Resource.parse[FooClass]("/test.yml") match { + case Some(foo) => + foo.id should equal("55555555") + case _ => fail() + } + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassFieldTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassFieldTest.scala new file mode 100644 index 0000000000..9eb811b457 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassFieldTest.scala @@ -0,0 +1,226 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.annotation.{JsonIgnore, JsonProperty} +import com.fasterxml.jackson.databind.JavaType +import com.fasterxml.jackson.databind.annotation.JsonDeserialize +import com.fasterxml.jackson.databind.deser.std.NumberDeserializers.BigDecimalDeserializer +import com.twitter.util.jackson.{ + JacksonScalaObjectMapperType, + ScalaObjectMapper, + TestInjectableValue, + _ +} +import com.twitter.util.validation.ScalaValidator +import jakarta.validation.constraints.NotEmpty +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class CaseClassFieldTest extends AnyFunSuite with Matchers { + private[this] val mapper: JacksonScalaObjectMapperType = + ScalaObjectMapper().underlying + + test("CaseClassField.createFields have field name foo") { + val fields = deserializerFor(classOf[WithEmptyJsonProperty]).fields + + fields.length should equal(1) + fields.head.name should equal("foo") + } + + test("CaseClassField.createFields also have field name foo") { + val fields = deserializerFor(classOf[WithoutJsonPropertyAnnotation]).fields + + fields.length should equal(1) + fields.head.name should equal("foo") + } + + test("CaseClassField.createFields have field name bar") { + val fields = deserializerFor(classOf[WithNonemptyJsonProperty]).fields + + fields.length should equal(1) + fields.head.name should equal("bar") + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation") { + val fields = deserializerFor(classOf[Aum]).fields + + fields.length should equal(2) + + val iField = fields.head + iField.name should equal("i") + iField.annotations.length should equal(1) + iField.annotations.head.annotationType() should be(classOf[JsonProperty]) + + val jField = fields.last + jField.name should equal("j") + jField.annotations.length should equal(1) + jField.annotations.head.annotationType() should be(classOf[JsonProperty]) + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 2") { + val fields = deserializerFor(classOf[FooBar]).fields + + fields.length should equal(1) + + val helloField = fields.head + helloField.name should equal("helloWorld") + helloField.annotations.length should equal(2) + helloField.annotations.head.annotationType() should be(classOf[JsonProperty]) + helloField.annotations.exists(_.annotationType() == classOf[TestInjectableValue]) should be( + true) + helloField.annotations.last + .asInstanceOf[TestInjectableValue].value() should be("accept") // from Bar + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 3") { + val fields = deserializerFor(classOf[TestTraitImpl]).fields + + /* + in trait: + --------- + @JsonProperty("oldness") + def age: Int + @NotEmpty + def name: String + + case class constructor: + ----------------------- + @JsonProperty("ageness") age: Int, // should override inherited annotation from trait + @TestInjectableValue name: String, // should have two annotations, one from trait and one here + @TestInjectableValue dateTime: LocalDate, + @JsonProperty foo: String, // empty JsonProperty should default to field name + @JsonDeserialize(contentAs = classOf[BigDecimal], using = classOf[BigDecimalDeserializer]) + double: BigDecimal, + @JsonIgnore ignoreMe: String + */ + fields.length should equal(6) + + val fieldMap: Map[String, CaseClassField] = + fields.map(field => field.name -> field).toMap + + val ageField = fieldMap("ageness") + ageField.annotations.length should equal(1) + ageField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + ageField.annotations.head.asInstanceOf[JsonProperty].value() should equal("ageness") + + val nameField = fieldMap("name") + nameField.annotations.length should equal(2) + nameField.annotations.exists(_.annotationType() == classOf[NotEmpty]) should be(true) + nameField.annotations.exists(_.annotationType() == classOf[TestInjectableValue]) should be(true) + + val dateTimeField = fieldMap("dateTime") + dateTimeField.annotations.length should equal(1) + dateTimeField.annotations.exists(_.annotationType() == classOf[TestInjectableValue]) should be( + true) + + val fooField = fieldMap("foo") + fooField.annotations.length should equal(1) + fooField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + fooField.annotations.head.asInstanceOf[JsonProperty].value() should equal("") + + val doubleField = fieldMap("double") + doubleField.annotations.length should equal(1) + doubleField.annotations.exists(_.annotationType() == classOf[JsonDeserialize]) should be(true) + doubleField.annotations.head.asInstanceOf[JsonDeserialize].contentAs() should be( + classOf[BigDecimal]) + doubleField.annotations.head.asInstanceOf[JsonDeserialize].using() should be( + classOf[BigDecimalDeserializer]) + + val ignoreMeField = fieldMap("ignoreMe") + ignoreMeField.annotations.length should equal(1) + ignoreMeField.annotations.exists(_.annotationType() == classOf[JsonIgnore]) should be(true) + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 4") { + val fields: Seq[CaseClassField] = deserializerFor(classOf[FooBaz]).fields + + fields.length should equal(1) + val helloField: CaseClassField = fields.head + helloField.annotations.length should equal(2) + helloField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + helloField.annotations.head + .asInstanceOf[JsonProperty].value() should equal("goodbyeWorld") // from Baz + + helloField.annotations.exists(_.annotationType() == classOf[TestInjectableValue]) should be( + true) + helloField.annotations.last + .asInstanceOf[TestInjectableValue].value() should be("accept") // from Bar + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 5") { + val fields = deserializerFor(classOf[FooBarBaz]).fields + + fields.length should equal(1) + val helloField: CaseClassField = fields.head + helloField.annotations.length should equal(2) + helloField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + helloField.annotations.head + .asInstanceOf[JsonProperty].value() should equal("goodbye") // from BarBaz + + helloField.annotations.exists(_.annotationType() == classOf[TestInjectableValue]) should be( + true) + helloField.annotations.last + .asInstanceOf[TestInjectableValue].value() should be("accept") // from Bar + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 6") { + val fields = deserializerFor(classOf[File]).fields + + fields.length should equal(1) + val uriField: CaseClassField = fields.head + uriField.annotations.length should equal(1) + uriField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + uriField.annotations.head.asInstanceOf[JsonProperty].value() should equal("file") + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 7") { + val fields = deserializerFor(classOf[Folder]).fields + + fields.length should equal(1) + val uriField: CaseClassField = fields.head + uriField.annotations.length should equal(1) + uriField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + uriField.annotations.head.asInstanceOf[JsonProperty].value() should equal("folder") + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 8") { + val fields = deserializerFor(classOf[LoadableFile]).fields + + fields.length should equal(1) + val uriField: CaseClassField = fields.head + uriField.annotations.length should equal(1) + uriField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + uriField.annotations.head.asInstanceOf[JsonProperty].value() should equal("file") + } + + test("CaseClassField.createFields sees inherited JsonProperty annotation 9") { + val fields = deserializerFor(classOf[LoadableFolder]).fields + + fields.length should equal(1) + val uriField: CaseClassField = fields.head + uriField.annotations.length should equal(1) + uriField.annotations.exists(_.annotationType() == classOf[JsonProperty]) should be(true) + uriField.annotations.head.asInstanceOf[JsonProperty].value() should equal("folder") + } + + test("Seq[Long]") { + val fields = deserializerFor(classOf[CaseClassWithArrayLong]).fields + + fields.length should equal(1) + val arrayField: CaseClassField = fields.head + arrayField.javaType.getTypeName should be( + "[array type, component type: [simple type, class long]]") + } + + private[this] def deserializerFor(clazz: Class[_]): CaseClassDeserializer = { + val javaType: JavaType = mapper.constructType(clazz) + new CaseClassDeserializer( + javaType = javaType, + mapper.getDeserializationConfig, + mapper.getDeserializationConfig.introspect(javaType), + Some(ScalaValidator()) + ) + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassValidationsTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassValidationsTest.scala new file mode 100644 index 0000000000..592567c887 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/CaseClassValidationsTest.scala @@ -0,0 +1,501 @@ +package com.twitter.util.jackson.caseclass + +import com.twitter.util.jackson._ +import com.twitter.util.jackson.caseclass.exceptions.CaseClassFieldMappingException.ValidationError +import com.twitter.util.jackson.caseclass.exceptions.{ + CaseClassFieldMappingException, + CaseClassMappingException +} +import com.twitter.util.validation.conversions.PathOps._ +import jakarta.validation.Path.ParameterNode +import jakarta.validation.{ElementKind, UnexpectedTypeException} +import java.time.format.DateTimeFormatter +import java.time.{LocalDate, LocalDateTime} +import org.junit.runner.RunWith +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +private object CaseClassValidationsTest { + val now: LocalDateTime = + LocalDateTime.parse("2015-04-09T05:17:15Z", DateTimeFormatter.ISO_OFFSET_DATE_TIME) + val BaseCar: Car = { + Car( + id = 1, + make = CarMake.Ford, + model = "Model-T", + year = 2000, + owners = Seq(), + numDoors = 2, + manual = true, + ownershipStart = now, + ownershipEnd = now.plusYears(15), + warrantyStart = Some(now), + warrantyEnd = Some(now.plusYears(5)), + passengers = Seq.empty[Person] + ) + } +} + +@RunWith(classOf[JUnitRunner]) +class CaseClassValidationsTest + extends AnyFunSuite + with BeforeAndAfterAll + with Matchers + with ScalaObjectMapperFunctions { + import CaseClassValidationsTest._ + + /* Class under test */ + override protected val mapper: ScalaObjectMapper = ScalaObjectMapper() + + test("deserialization#generic types") { + // passes validation + parse[GenericTestCaseClassWithValidation[String]]("""{"data" : "widget"}""") should equal( + GenericTestCaseClassWithValidation("widget")) + + val e1 = intercept[CaseClassMappingException] { + parse[GenericTestCaseClassWithValidationAndMultipleArgs[String]]( + """{"data" : "", "number" : 3}""") + } + e1.errors.size should equal(2) + val e2 = intercept[CaseClassMappingException] { + parse[GenericTestCaseClassWithValidationAndMultipleArgs[String]]( + """{"data" : "widget", "number" : 3}""") + } + e2.errors.size should equal(1) + // passes validation + parse[GenericTestCaseClassWithValidationAndMultipleArgs[String]]( + """{"data" : "widget", "number" : 5}""") + } + + test("enums#default validation") { + val e = intercept[CaseClassMappingException] { + parse[CaseClassWithNotEmptyValidation]("""{ + "name" : "", + "make" : "foo" + }""") + } + e.errors.map { _.getMessage } should equal( + Seq( + """make: 'foo' is not a valid CarMakeEnum with valid values: ford, vw""", + "name: must not be empty") + ) + } + + test("enums#invalid validation") { + // the name field is annotated with @ThrowsRuntimeExceptionConstraint which is implemented + // by the ThrowsRuntimeExceptionConstraintValidator...that simply throws a RuntimeException + val e = intercept[RuntimeException] { + parse[CaseClassWithInvalidValidation]("""{ + "name" : "Bob", + "make" : "foo" + }""") + } + + e.getMessage should be("validator foo error") + } + + test("enums#invalid validation 1") { + // the name field is annotated with @NoValidatorImplConstraint which has no constraint validator + // specified nor registered and thus errors when applied. + val e = intercept[UnexpectedTypeException] { + parse[CaseClassWithNoValidatorImplConstraint]("""{ + "name" : "Bob", + "make" : "foo" + }""") + } + + e.getMessage.contains( + "No validator could be found for constraint 'interface com.twitter.util.jackson.NoValidatorImplConstraint' validating type 'java.lang.String'. Check configuration for 'CaseClassWithNoValidatorImplConstraint.name'") should be( + true) + } + + test("method validation exceptions") { + val e = intercept[RuntimeException] { + parse[CaseClassWithMethodValidationException]( + """{ + | "name" : "foo barr", + | "passphrase" : "the quick brown fox jumps over the lazy dog" + |}""".stripMargin) + } + e.getMessage.contains("oh noes") + } + + /* + Notes: + + ValidationError.Field types will not have a root bean instance and the violation property path + will include the class name, e.g., for a violation on the `id` parameter of the `RentalStation` case + class, the property path would be: `RentalStation.id`. + + ValidationError.Method types will have a root bean instance (since they require an instance in order for + the method to be invoked) and the violation property path will not include the class name since this + is understood to be known since the method validations typically have to be manually invoked, e.g., + for the method `validateState` that is annotated with field `state`, the property path of a violation would be: + `validateState.state`. + */ + + test("class and field level validations#success") { + parse[Car](BaseCar) + } + + test("class and field level validations#top-level failed validations") { + val value = BaseCar.copy(id = 2, year = 1910) + val e = intercept[CaseClassMappingException] { + parse[Car](value) + } + + e.errors.size should equal(1) + + val error = e.errors.head + error.path should equal(CaseClassFieldMappingException.PropertyPath.leaf("year")) + error.reason.message should equal("must be greater than or equal to 2000") + error.reason.detail match { + case ValidationError(violation, ValidationError.Field, None) => + violation.getPropertyPath.toString should equal("Car.year") + violation.getMessage should equal("must be greater than or equal to 2000") + violation.getInvalidValue should equal(1910) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean == null should be( + true + ) // ValidationError.Field types won't have a root bean instance + case _ => fail() + } + } + + test("method validation#top-level") { + val address = Address( + street = Some("123 Main St."), + city = "Tampa", + state = "FL" // invalid + ) + + val e = intercept[CaseClassMappingException] { + parse[Address](address) + } + + e.errors.size should equal(1) + val errors = e.errors + + errors.head.path should equal(CaseClassFieldMappingException.PropertyPath.Empty) + errors.head.reason.message should equal("state must be one of [CA, MD, WI]") + errors.head.reason.detail.getClass should equal(classOf[ValidationError]) + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal("state must be one of [CA, MD, WI]") + violation.getInvalidValue should equal(address) + violation.getRootBeanClass should equal(classOf[Address]) + violation.getRootBean should equal( + address + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "validateState" + ) // the @MethodValidation annotation does not specify a field. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PROPERTY) + leafNode.getName should equal("validateState") + case _ => + fail() + } + } + + test("class and field level validations#nested failed validations") { + val invalidAddress = + Address( + street = Some(""), // invalid + city = "", // invalid + state = "FL" // invalid + ) + + val owners = Seq( + Person( + id = 1, + name = "joe smith", + dob = Some(LocalDate.now), + age = None, + address = Some(invalidAddress) + ) + ) + val car = BaseCar.copy(owners = owners) + + val e = intercept[CaseClassMappingException] { + parse[Car](car) + } + + e.errors.size should equal(2) + val errors = e.errors + + errors.head.path should equal( + CaseClassFieldMappingException.PropertyPath + .leaf("city").withParent("address").withParent("owners")) + errors.head.reason.message should equal("must not be empty") + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Field, None) => + violation.getPropertyPath.toString should equal("Address.city") + violation.getMessage should equal("must not be empty") + violation.getInvalidValue should equal("") + violation.getRootBeanClass should equal(classOf[Address]) + violation.getRootBean == null should be( + true + ) // ValidationError.Field types won't have a root bean instance + case _ => + fail() + } + + errors(1).path should equal( + CaseClassFieldMappingException.PropertyPath + .leaf("street").withParent("address").withParent("owners") + ) + errors(1).reason.message should equal("must not be empty") + errors(1).reason.detail match { + case ValidationError(violation, ValidationError.Field, None) => + violation.getPropertyPath.toString should equal("Address.street") + violation.getMessage should equal("must not be empty") + violation.getInvalidValue should equal("") + violation.getRootBeanClass should equal(classOf[Address]) + violation.getRootBean == null should be( + true + ) // ValidationError.Field types won't have a root bean instance + case _ => + fail() + } + } + + test("class and field level validations#nested method validations") { + val owners = Seq( + Person( + id = 2, + name = "joe smith", + dob = Some(LocalDate.now), + age = None, + address = Some(Address(city = "pyongyang", state = "KP" /* invalid */ )) + ) + ) + val car: Car = BaseCar.copy(owners = owners) + + val e = intercept[CaseClassMappingException] { + parse[Car](car) + } + + e.errors.size should equal(1) + val errors = e.errors + + errors.head.path should equal( + CaseClassFieldMappingException.PropertyPath + .leaf("address").withParent("owners") + ) + errors.head.reason.message should equal("state must be one of [CA, MD, WI]") + errors.head.reason.detail.getClass should equal(classOf[ValidationError]) + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal("state must be one of [CA, MD, WI]") + violation.getInvalidValue should equal(Address(city = "pyongyang", state = "KP")) + violation.getRootBeanClass should equal(classOf[Address]) + violation.getRootBean should equal( + Address(city = "pyongyang", state = "KP") + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "validateState" + ) // the @MethodValidation annotation does not specify a field. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PROPERTY) + leafNode.getName should equal("validateState") + case _ => + fail() + } + + errors.map(_.getMessage) should equal(Seq("owners.address: state must be one of [CA, MD, WI]")) + } + + test("class and field level validations#end before start") { + val value: Car = + BaseCar.copy(ownershipStart = BaseCar.ownershipEnd, ownershipEnd = BaseCar.ownershipStart) + val e = intercept[CaseClassMappingException] { + parse[Car](value) + } + + e.errors.size should equal(1) + val errors = e.errors + + errors.head.path should equal(CaseClassFieldMappingException.PropertyPath.leaf("ownership_end")) + errors.head.reason.message should equal( + "ownershipEnd [2015-04-09T05:17:15] must be after ownershipStart [2030-04-09T05:17:15]") + errors.head.reason.detail.getClass should equal(classOf[ValidationError]) + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal( + "ownershipEnd [2015-04-09T05:17:15] must be after ownershipStart [2030-04-09T05:17:15]") + violation.getInvalidValue should equal(value) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean should equal( + value + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "ownershipTimesValid.ownershipEnd" + ) // the @MethodValidation annotation specifies a single field. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PARAMETER) + leafNode.getName should equal("ownershipEnd") + leafNode.as(classOf[ParameterNode]) // shouldn't blow up + case _ => + fail() + } + } + + test("class and field level validations#optional end before start") { + val value: Car = + BaseCar.copy(warrantyStart = BaseCar.warrantyEnd, warrantyEnd = BaseCar.warrantyStart) + val e = intercept[CaseClassMappingException] { + parse[Car](value) + } + + e.errors.size should equal(2) + val errors = e.errors + + errors.head.path should equal(CaseClassFieldMappingException.PropertyPath.leaf("warranty_end")) + errors.head.reason.message should equal( + "warrantyEnd [2015-04-09T05:17:15] must be after warrantyStart [2020-04-09T05:17:15]" + ) + errors.head.reason.detail.getClass should equal(classOf[ValidationError]) + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal( + "warrantyEnd [2015-04-09T05:17:15] must be after warrantyStart [2020-04-09T05:17:15]") + violation.getInvalidValue should equal(value) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean should equal( + value + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "warrantyTimeValid.warrantyEnd" + ) // the @MethodValidation annotation specifies two fields. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PARAMETER) + leafNode.getName should equal("warrantyEnd") + leafNode.as(classOf[ParameterNode]) // shouldn't blow up + case _ => + fail() + } + + errors(1).path should equal(CaseClassFieldMappingException.PropertyPath.leaf("warranty_start")) + errors(1).reason.message should equal( + "warrantyEnd [2015-04-09T05:17:15] must be after warrantyStart [2020-04-09T05:17:15]" + ) + errors(1).reason.detail.getClass should equal(classOf[ValidationError]) + errors(1).reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal( + "warrantyEnd [2015-04-09T05:17:15] must be after warrantyStart [2020-04-09T05:17:15]") + violation.getInvalidValue should equal(value) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean should equal( + value + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "warrantyTimeValid.warrantyStart" + ) // the @MethodValidation annotation specifies two fields. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PARAMETER) + leafNode.getName should equal("warrantyStart") + leafNode.as(classOf[ParameterNode]) // shouldn't blow up + case _ => + fail() + } + } + + test("class and field level validations#no start with end") { + val value: Car = BaseCar.copy(warrantyStart = None, warrantyEnd = BaseCar.warrantyEnd) + val e = intercept[CaseClassMappingException] { + parse[Car](value) + } + + e.errors.size should equal(2) + val errors = e.errors + + errors.head.path should equal(CaseClassFieldMappingException.PropertyPath.leaf("warranty_end")) + errors.head.reason.message should equal( + "both warrantyStart and warrantyEnd are required for a valid range" + ) + errors.head.reason.detail.getClass should equal(classOf[ValidationError]) + errors.head.reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal( + "both warrantyStart and warrantyEnd are required for a valid range") + violation.getInvalidValue should equal(value) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean should equal( + value + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "warrantyTimeValid.warrantyEnd" + ) // the @MethodValidation annotation specifies two fields. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PARAMETER) + leafNode.getName should equal("warrantyEnd") + leafNode.as(classOf[ParameterNode]) // shouldn't blow up + case _ => + fail() + } + + errors(1).path should equal(CaseClassFieldMappingException.PropertyPath.leaf("warranty_start")) + errors(1).reason.message should equal( + "both warrantyStart and warrantyEnd are required for a valid range" + ) + errors(1).reason.detail.getClass should equal(classOf[ValidationError]) + errors(1).reason.detail match { + case ValidationError(violation, ValidationError.Method, None) => + violation.getMessage should equal( + "both warrantyStart and warrantyEnd are required for a valid range") + violation.getInvalidValue should equal(value) + violation.getRootBeanClass should equal(classOf[Car]) + violation.getRootBean should equal( + value + ) // ValidationError.Method types have a root bean instance + violation.getPropertyPath.toString should equal( + "warrantyTimeValid.warrantyStart" + ) // the @MethodValidation annotation specifies two fields. + val leafNode = violation.getPropertyPath.getLeafNode + leafNode.getKind should equal(ElementKind.PARAMETER) + leafNode.getName should equal("warrantyStart") + leafNode.as(classOf[ParameterNode]) // shouldn't blow up + case _ => + fail() + } + } + + test("class and field level validations#errors sorted by message") { + val first = + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.Empty, + CaseClassFieldMappingException.Reason("123")) + val second = + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.Empty, + CaseClassFieldMappingException.Reason("aaa")) + val third = CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.leaf("bla"), + CaseClassFieldMappingException.Reason("zzz")) + val fourth = + CaseClassFieldMappingException( + CaseClassFieldMappingException.PropertyPath.Empty, + CaseClassFieldMappingException.Reason("xxx")) + + val unsorted = Set(third, second, fourth, first) + val expectedSorted = Seq(first, second, third, fourth) + + CaseClassMappingException(unsorted).errors should equal((expectedSorted)) + } + + test("class and field level validations#option[string] validation") { + val address = Address( + street = Some(""), // invalid + city = "New Orleans", + state = "LA" + ) + + intercept[CaseClassMappingException] { + mapper.parse[Address](mapper.writeValueAsBytes(address)) + } + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/OptionalValidationTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/OptionalValidationTest.scala new file mode 100644 index 0000000000..aaf134cb43 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/OptionalValidationTest.scala @@ -0,0 +1,87 @@ +package com.twitter.util.jackson.caseclass + +import com.twitter.util.WrappedValue +import com.twitter.util.jackson.ScalaObjectMapper +import com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException +import com.twitter.util.validation.MethodValidation +import com.twitter.util.validation.constraints.OneOf +import com.twitter.util.validation.engine.MethodValidationResult +import jakarta.validation.constraints.{Min, NotEmpty} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +private object OptionalValidationTest { + case class State( + @OneOf(Array("active", "inactive")) + state: String) + extends WrappedValue[String] + + case class Threshold( + @NotEmpty id: Option[String], + @Min(0) lowerBound: Int, + @Min(0) upperBound: Int, + state: State) { + @MethodValidation + def method: MethodValidationResult = MethodValidationResult.validIfTrue( + lowerBound <= upperBound, + "Lower Bound cannot be greater than Upper Bound" + ) + } +} + +@RunWith(classOf[JUnitRunner]) +class OptionalValidationTest extends AnyFunSuite with Matchers { + import OptionalValidationTest._ + + private val defaultMapper = + ScalaObjectMapper.builder.objectMapper + + private val nullValidationMapper = + ScalaObjectMapper.builder.withNoValidation.objectMapper + + test("default mapper will trigger NotEmpty validation") { + val invalid = Threshold(Some(""), 1, 4, State("active")) + intercept[CaseClassMappingException] { + check(defaultMapper, invalid) + } + } + + test("default mapper will trigger lower and upper min validation") { + val invalid = Threshold(Some(""), -1, -1, State("active")) + intercept[CaseClassMappingException] { + check(defaultMapper, invalid) + } + } + + test("default mapper will trigger invalid state") { + val invalid = Threshold(Some("id"), 4, 6, State("other")) + intercept[CaseClassMappingException] { + check(defaultMapper, invalid) + } + } + + test("default mapper will trigger method validation") { + val invalid = Threshold(None, 8, 4, State("active")) + intercept[CaseClassMappingException] { + check(defaultMapper, invalid) + } + } + + test("no-op mapper will allow values that are considered invalid") { + val invalid = Threshold(Some(""), -1, -3, State("other")) + check(nullValidationMapper, invalid) + } + + test("no-op mapper will never call method validation") { + val invalid = Threshold(None, 8, 4, State("active")) + check(nullValidationMapper, invalid) + } + + def check[T: Manifest](mapper: ScalaObjectMapper, check: T): Unit = { + val str = mapper.writeValueAsString(check) + val result = mapper.parse[T](str) + check should equal(result) + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/RichJsonProcessingExceptionTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/RichJsonProcessingExceptionTest.scala new file mode 100644 index 0000000000..7afade0e5c --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/RichJsonProcessingExceptionTest.scala @@ -0,0 +1,92 @@ +package com.twitter.util.jackson.caseclass + +import com.fasterxml.jackson.core.{JsonGenerator, JsonParseException, JsonParser} +import com.fasterxml.jackson.databind.JsonMappingException +import com.fasterxml.jackson.databind.`type`.SimpleType +import com.fasterxml.jackson.databind.exc.{ + IgnoredPropertyException, + InvalidDefinitionException, + MismatchedInputException +} +import com.twitter.util.jackson.CaseClassWithAllTypes +import com.twitter.util.jackson.caseclass.exceptions._ +import com.twitter.util.mock.Mockito +import java.util +import java.util.Collections +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class RichJsonProcessingExceptionTest extends AnyFunSuite with Matchers with Mockito { + private[this] val genericType = SimpleType.constructUnsafe(classOf[String]) + private val message = "this is a test" + + test("JsonMappingException with cause") { + val cause = new IllegalArgumentException(message) + val e = new JsonMappingException(null, message, cause) + + // should be the mapping exception message + e.errorMessage should equal(message) + } + + test("JsonMappingException with no cause") { + val e = new JsonMappingException(null, message) + + // should be the mapping exception message + e.errorMessage should equal(message) + } + + test("MismatchedInputException with cause") { + val mockJsonParser = mock[JsonParser] + + val e = MismatchedInputException.from( + mockJsonParser, + genericType, + message + ) + + // should be from the cause + e.errorMessage should equal(message) + } + + test("InvalidDefinitionException with no cause") { + val mockJsonGenerator = mock[JsonGenerator] + + val e = InvalidDefinitionException.from(mockJsonGenerator, message, genericType) + + // should be from the invalid definition exception message via `getOriginalMessage` + e.errorMessage should equal(message) + } + + test("IgnoredPropertyException with no cause") { + val mockJsonParser = mock[JsonParser] + + val e = IgnoredPropertyException.from( + mockJsonParser, + CaseClassWithAllTypes.getClass, + "noname", + Collections.emptySet().asInstanceOf[util.Collection[Object]]) + + // should be from the ignored property exception message via `getOriginalMessage` + e.errorMessage should equal( + """Ignored field "noname" (class com.twitter.util.jackson.CaseClassWithAllTypes$) encountered; mapper configured not to allow this""") + } + + test("JsonParseException with cause") { + val cause = new IllegalArgumentException(message) + val e = new JsonParseException(null, "this is not the message", cause) + + // should be from the cause + e.errorMessage should equal(message) + } + + test("JsonParseException with no cause") { + val e = new JsonParseException(null, message) + + // should be from the parse exception message via `getOriginalMessage` + e.errorMessage should equal(message) + } + +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/TypesTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/TypesTest.scala new file mode 100644 index 0000000000..efa1656de8 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclass/TypesTest.scala @@ -0,0 +1,34 @@ +package com.twitter.util.jackson.caseclass + +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class TypesTest extends AnyFunSuite with Matchers { + + test("types#handles null") { + val clazz = Types.wrapperType(null) + clazz != null should be(true) + clazz.getSimpleName should equal(classOf[Null].getSimpleName) + } + + test("types#all types") { + classOf[Null].isAssignableFrom(Types.wrapperType(null)) should be(true) + + classOf[java.lang.Byte].isAssignableFrom(Types.wrapperType(java.lang.Byte.TYPE)) should be(true) + classOf[java.lang.Short] + .isAssignableFrom(Types.wrapperType(java.lang.Short.TYPE)) should be(true) + classOf[java.lang.Character] + .isAssignableFrom(Types.wrapperType(java.lang.Character.TYPE)) should be(true) + classOf[java.lang.Integer] + .isAssignableFrom(Types.wrapperType(java.lang.Integer.TYPE)) should be(true) + classOf[java.lang.Long].isAssignableFrom(Types.wrapperType(java.lang.Long.TYPE)) should be(true) + classOf[java.lang.Double] + .isAssignableFrom(Types.wrapperType(java.lang.Double.TYPE)) should be(true) + classOf[java.lang.Boolean] + .isAssignableFrom(Types.wrapperType(java.lang.Boolean.TYPE)) should be(true) + classOf[java.lang.Void].isAssignableFrom(Types.wrapperType(java.lang.Void.TYPE)) should be(true) + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/caseclasses.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclasses.scala new file mode 100644 index 0000000000..2382f358e7 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/caseclasses.scala @@ -0,0 +1,847 @@ +package com.twitter.util.jackson + +import com.fasterxml.jackson.annotation.JsonSubTypes.Type +import com.fasterxml.jackson.annotation._ +import com.fasterxml.jackson.core.JsonParser +import com.fasterxml.jackson.core.`type`.TypeReference +import com.fasterxml.jackson.databind._ +import com.fasterxml.jackson.databind.annotation.JsonDeserialize +import com.fasterxml.jackson.databind.annotation.JsonNaming +import com.fasterxml.jackson.databind.deser.std.NumberDeserializers.BigDecimalDeserializer +import com.fasterxml.jackson.databind.deser.std.NumberDeserializers.NumberDeserializer +import com.fasterxml.jackson.databind.deser.std.StdDeserializer +import com.fasterxml.jackson.databind.node.ValueNode +import com.fasterxml.jackson.module.scala.JsonScalaEnumeration +import com.twitter.util.Time +import com.twitter.util.WrappedValue +import com.twitter.util.jackson.caseclass.SerdeLogging +import com.twitter.util.validation.MethodValidation +import com.twitter.util.validation.constraints.OneOf +import com.twitter.util.validation.constraints.UUID +import com.twitter.util.validation.engine.MethodValidationResult +import com.twitter.{util => ctu} +import jakarta.validation.Payload +import jakarta.validation.constraints._ +import java.time.format.DateTimeFormatter +import java.time.temporal.Temporal +import java.time.LocalDate +import java.time.LocalDateTime +import scala.math.BigDecimal.RoundingMode + +@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "type") +@JsonSubTypes( + Array( + new Type(value = classOf[BoundaryA], name = "a"), + new Type(value = classOf[BoundaryB], name = "b"), + new Type(value = classOf[BoundaryC], name = "c") + ) +) +trait Boundary { + def value: String +} +case class BoundaryA(value: String) extends Boundary +case class BoundaryB(value: String) extends Boundary +case class BoundaryC(value: String) extends Boundary + +case class GG[T, U, V, X, Y, Z](t: T, u: U, v: V, x: X, y: Y, z: Z) +case class B(value: String) +case class D(value: Int) +case class AGeneric[T](b: Option[T]) +case class AAGeneric[T, U](b: Option[T], c: Option[U]) +case class C(a: AGeneric[B]) +case class E[T, U, V](a: AAGeneric[T, U], b: Option[V]) +case class F[T, U, V, X, Y, Z](a: Option[T], b: Seq[U], c: Option[V], d: X, e: Either[Y, Z]) +case class WithTypeBounds[A <: Boundary](a: Option[A]) +case class G[T, U, V, X, Y, Z](gg: GG[T, U, V, X, Y, Z]) + +case class SomeNumberType[N <: Number]( + @JsonDeserialize(contentAs = classOf[Number], using = classOf[NumberDeserializer]) n: Option[N]) + +object Weekday extends Enumeration { + type Weekday = Value + val Mon, Tue, Wed, Thu, Fri, Sat, Sun = Value +} +class WeekdayType extends TypeReference[Weekday.type] +object Month extends Enumeration { + type Month = Value + val Jan, Feb, Mar, Apr, May, Jun, Jul, Aug, Sep, Oct, Nov, Dec = Value +} +class MonthType extends TypeReference[Month.type] +case class BasicDate( + @JsonScalaEnumeration(classOf[MonthType]) month: Month.Month, + day: Int, + year: Int, + @JsonScalaEnumeration(classOf[WeekdayType]) weekday: Weekday.Weekday) + +case class WithOptionalScalaEnumeration( + @JsonScalaEnumeration(classOf[MonthType]) month: Option[Month.Month]) + +@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "type") +@JsonSubTypes( + Array( + new Type(value = classOf[Rectangle], name = "rectangle"), + new Type(value = classOf[Circle], name = "circle") + ) +) +sealed trait Shape +case class Rectangle(@Min(0) width: Int, @Min(0) height: Int) extends Shape +case class Circle(@Min(0) radius: Int) extends Shape +case class View(shapes: Seq[Shape]) +case class OptionalView(shapes: Seq[Shape], optional: Option[Shape]) + +object TestJsonCreator { + @JsonCreator + def apply(s: String): TestJsonCreator = TestJsonCreator(s.toInt) +} +case class TestJsonCreator(int: Int) + +object TestJsonCreator2 { + @JsonCreator + def apply(strings: Seq[String]): TestJsonCreator2 = TestJsonCreator2(strings.map(_.toInt)) +} +case class TestJsonCreator2(ints: Seq[Int], default: String = "Hello, World") + +object TestJsonCreatorWithValidation { + @JsonCreator + def apply(@NotEmpty s: String): TestJsonCreatorWithValidation = + TestJsonCreatorWithValidation(s.toInt) +} +case class TestJsonCreatorWithValidation(int: Int) + +object TestJsonCreatorWithValidations { + @JsonCreator + def apply(@NotEmpty @OneOf(Array("42", "137")) s: String): TestJsonCreatorWithValidations = + TestJsonCreatorWithValidations(s.toInt) +} +case class TestJsonCreatorWithValidations(int: Int) + +case class CaseClassWithMultipleConstructors(number1: Long, number2: Long, number3: Long) { + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } +} + +case class CaseClassWithMultipleConstructorsAnnotated(number1: Long, number2: Long, number3: Long) { + @JsonCreator + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } +} + +case class CaseClassWithMultipleConstructorsAnnotatedAndValidations( + number1: Long, + number2: Long, + uuid: String) { + @JsonCreator + def this( + @NotEmpty numberAsString1: String, + @OneOf(Array("10001", "20002", "30003")) numberAsString2: String, + @UUID thirdArgument: String + ) { + this(numberAsString1.toLong, numberAsString2.toLong, thirdArgument) + } +} + +case class TimeWithFormat(@JsonFormat(pattern = "yyyy-MM-dd'T'HH:mm:ss.SSSXXX") when: Time) + +/* Note: the decoder automatically changes "_i" to "i" for de/serialization: + * See CaseClassField#jsonNameForField */ +trait Aumly { @JsonProperty("i") def _i: Int; @JsonProperty("j") def _j: String } +case class Aum(_i: Int, _j: String) extends Aumly + +trait Bar { + @JsonProperty("helloWorld") @TestInjectableValue(value = "accept") + def hello: String +} +case class FooBar(hello: String) extends Bar + +trait Baz extends Bar { + @JsonProperty("goodbyeWorld") + def hello: String +} +case class FooBaz(hello: String) extends Baz + +trait BarBaz { + @JsonProperty("goodbye") + def hello: String +} +case class FooBarBaz(hello: String) + extends BarBaz + with Bar // will end up with BarBaz @JsonProperty value as trait linearization is "right-to-left" + +trait Loadable { + @JsonProperty("url") + def uri: String +} +abstract class Resource { + @JsonProperty("resource") + def uri: String +} +case class File(@JsonProperty("file") uri: String) extends Resource +case class Folder(@JsonProperty("folder") uri: String) extends Resource + +abstract class LoadableResource extends Loadable { + @JsonProperty("resource") + override def uri: String +} +case class LoadableFile(@JsonProperty("file") uri: String) extends LoadableResource +case class LoadableFolder(@JsonProperty("folder") uri: String) extends LoadableResource + +trait TestTrait { + @JsonProperty("oldness") + def age: Int + @NotEmpty + def name: String +} +@JsonNaming +case class TestTraitImpl( + @JsonProperty("ageness") age: Int, // should override inherited annotation from trait + @TestInjectableValue name: String, // should have two annotations, one from trait and one here + @TestInjectableValue dateTime: LocalDate, + @JsonProperty foo: String, + @JsonDeserialize(contentAs = classOf[BigDecimal], using = classOf[BigDecimalDeserializer]) + double: BigDecimal, + @JsonIgnore ignoreMe: String) + extends TestTrait { + + lazy val testFoo: String = "foo" + lazy val testBar: String = "bar" +} + +sealed trait CarType { + @JsonValue + def toJson: String +} +object Volvo extends CarType { + override def toJson: String = "volvo" +} +object Audi extends CarType { + override def toJson: String = "audi" +} +object Volkswagen extends CarType { + override def toJson: String = "vw" +} + +case class Vehicle(vin: String, `type`: CarType) + +sealed trait ZeroOrOne +object Zero extends ZeroOrOne +object One extends ZeroOrOne +object Two extends ZeroOrOne + +case class CaseClassWithZeroOrOne(id: ZeroOrOne) + +case class CaseClass(id: Long, name: String) +case class CaseClassIdAndOption(id: Long, name: Option[String]) +@JsonIgnoreProperties(ignoreUnknown = false) +case class StrictCaseClassIdAndOption(id: Long, name: Option[String]) + +case class SimpleClassWithInjection(@TestInjectableValue(value = "accept") hello: String) + +@JsonIgnoreProperties(ignoreUnknown = false) +case class StrictCaseClass(id: Long, name: String) + +@JsonIgnoreProperties(ignoreUnknown = false) +case class StrictCaseClassWithOption(value: Option[String] = None) + +@JsonIgnoreProperties(ignoreUnknown = true) +case class LooseCaseClass(id: Long, name: String) + +case class CaseClassWithLazyVal(id: Long) { + lazy val woo = "yeah" +} + +case class GenericTestCaseClass[T](data: T) +case class GenericTestCaseClassWithValidation[T](@NotEmpty data: T) +case class GenericTestCaseClassWithValidationAndMultipleArgs[T]( + @NotEmpty data: T, + @Min(5) number: Int) + +case class Page[T](data: List[T], pageSize: Int, next: Option[Long], previous: Option[Long]) + +case class CaseClassWithGeneric[T](inside: GenericTestCaseClass[T]) + +case class CaseClassWithOptionalGeneric[T](inside: Option[GenericTestCaseClass[T]]) + +case class CaseClassWithTypes[T, U](first: T, second: U) + +case class CaseClassWithMapTypes[T, U](data: Map[T, U]) + +case class CaseClassWithManyTypes[R, S, T](one: R, two: S, three: T) + +case class CaseClassIgnoredFieldInConstructorNoDefault( + @JsonIgnore id: Long, + name: String, + description: String) + +case class CaseClassIgnoredFieldInConstructorWithDefault( + @JsonIgnore id: Long = 42L, + name: String, + description: String) + +case class CaseClassWithIgnoredField(id: Long) { + @JsonIgnore + val ignoreMe = "Foo" +} + +@JsonIgnoreProperties(Array("ignore_me", "feh")) +case class CaseClassWithIgnoredFieldsMatchAfterToSnakeCase(id: Long) { + val ignoreMe = "Foo" + val feh = "blah" +} + +@JsonIgnoreProperties(Array("ignore_me", "feh")) +case class CaseClassWithIgnoredFieldsExactMatch(id: Long) { + val ignore_me = "Foo" + val feh = "blah" +} + +case class CaseClassWithTransientField(id: Long) { + @transient + val lol = "asdf" +} + +case class CaseClassWithLazyField(id: Long) { + @JsonIgnore lazy val lol = "asdf" +} + +case class CaseClassWithOverloadedField(id: Long) { + def id(prefix: String): String = prefix + id +} + +case class CaseClassWithOption(value: Option[String] = None) + +case class CaseClassWithOptionAndValidation(@Min(100) towingCapacity: Option[Int] = None) + +case class CaseClassWithOptionAndIncompatibleValidation(@Min(100) areYou: Option[Boolean] = None) + +case class CaseClassWithJsonNode(value: JsonNode) + +case class CaseClassWithAllTypes( + map: Map[String, String], + set: Set[Int], + string: String, + list: List[Int], + seq: Seq[Int], + indexedSeq: IndexedSeq[Int], + vector: Vector[Int], + bigDecimal: BigDecimal, + bigInt: Int, //TODO: BigInt, + int: Int, + long: Long, + char: Char, + bool: Boolean, + short: Short, + byte: Byte, + float: Float, + double: Double, + any: Any, + anyRef: AnyRef, + intMap: Map[Int, Int] = Map(), + longMap: Map[Long, Long] = Map()) + +case class CaseClassWithException() { + throw JsonMappingException.from(null.asInstanceOf[JsonParser], "Oops!!!") +} + +object OuterObject { + + case class NestedCaseClass(id: Long) + + object InnerObject { + + case class SuperNestedCaseClass(id: Long) + } +} + +case class CaseClassWithSnakeCase(oneThing: String, twoThing: String) + +case class CaseClassWithArrays( + one: String, + two: Array[String], + three: Array[Int], + four: Array[Long], + five: Array[Char], + bools: Array[Boolean], + bytes: Array[Byte], + doubles: Array[Double], + floats: Array[Float]) + +case class CaseClassWithArrayLong(array: Array[Long]) + +case class CaseClassWithArrayListOfIntegers(arraylist: java.util.ArrayList[java.lang.Integer]) + +case class CaseClassWithArrayBoolean(array: Array[Boolean]) + +case class CaseClassWithArrayWrappedValueLong(array: Array[WrappedValueLong]) + +case class CaseClassWithSeqLong(seq: Seq[Long]) + +case class CaseClassWithSeqWrappedValueLong(seq: Seq[WrappedValueLong]) + +case class CaseClassWithValidation(@Min(1) value: Long) + +case class CaseClassWithSeqOfCaseClassWithValidation(seq: Seq[CaseClassWithValidation]) + +case class WrappedValueLongWithValidation(@Min(1) value: Long) extends WrappedValue[Long] + +case class CaseClassWithSeqWrappedValueLongWithValidation(seq: Seq[WrappedValueLongWithValidation]) + +case class Foo(name: String) + +case class CaseClassCharacter(c: Char) + +case class Car( + id: Long, + make: CarMake, + model: String, + @Min(2000) year: Int, + owners: Seq[Person], + @Min(0) numDoors: Int = 4, + manual: Boolean = false, + ownershipStart: LocalDateTime = LocalDateTime.now, + ownershipEnd: LocalDateTime = LocalDateTime.now.plusYears(1), + warrantyStart: Option[LocalDateTime] = None, + warrantyEnd: Option[LocalDateTime] = None, + passengers: Seq[Person] = Seq()) { + + @MethodValidation + def validateId: MethodValidationResult = { + MethodValidationResult.validIfTrue(id % 2 == 1, "id may not be even") + } + + @MethodValidation + def validateYearBeforeNow: MethodValidationResult = { + val thisYear = LocalDate.now.getYear + val yearMoreThanOneYearInFuture: Boolean = + if (year > thisYear) { (year - thisYear) > 1 } + else false + MethodValidationResult.validIfFalse( + yearMoreThanOneYearInFuture, + "Model year can be at most one year newer." + ) + } + + @MethodValidation(fields = Array("ownershipEnd")) + def ownershipTimesValid: MethodValidationResult = + validateTimeRange( + ownershipStart, + ownershipEnd, + "ownershipStart", + "ownershipEnd" + ) + + @MethodValidation(fields = Array("warrantyStart", "warrantyEnd")) + def warrantyTimeValid: MethodValidationResult = + validateTimeRange( + warrantyStart, + warrantyEnd, + "warrantyStart", + "warrantyEnd" + ) + + private[this] def validateTimeRange( + start: Option[LocalDateTime], + end: Option[LocalDateTime], + startProperty: String, + endProperty: String + ): MethodValidationResult = { + + val rangeDefined = start.isDefined && end.isDefined + val partialRange = !rangeDefined && (start.isDefined || end.isDefined) + + if (rangeDefined) + validateTimeRange(start.get, end.get, startProperty, endProperty) + else if (partialRange) + MethodValidationResult.Invalid( + "both %s and %s are required for a valid range".format(startProperty, endProperty) + ) + else + MethodValidationResult.Valid + } + + private[this] def validateTimeRange( + start: LocalDateTime, + end: LocalDateTime, + startProperty: String, + endProperty: String + ): MethodValidationResult = + MethodValidationResult.validIfTrue( + start.isBefore(end), + "%s [%s] must be after %s [%s]" + .format( + endProperty, + DateTimeFormatter.ISO_DATE_TIME.format(end), + startProperty, + DateTimeFormatter.ISO_DATE_TIME.format(start)) + ) +} + +case class PersonWithLogging( + id: Int, + name: String, + age: Option[Int], + age_with_default: Option[Int] = None, + nickname: String = "unknown") + extends SerdeLogging + +case class PersonWithDottedName(id: Int, @JsonProperty("name.last") lastName: String) + +case class SimplePerson(name: String) + +case class PersonWithThings( + id: Int, + name: String, + age: Option[Int], + @Size(min = 1, max = 10) things: Map[String, Things]) + +case class Things(@Size(min = 1, max = 2) names: Seq[String]) + +@JsonNaming +case class CamelCaseSimplePerson(myName: String) + +case class CamelCaseSimplePersonNoAnnotation(myName: String) + +case class CaseClassWithMap(map: Map[String, String]) + +case class CaseClassWithSortedMap(sortedMap: scala.collection.SortedMap[String, Int]) + +case class CaseClassWithSetOfLongs(set: Set[Long]) + +case class CaseClassWithSeqOfLongs(seq: Seq[Long]) + +case class CaseClassWithNestedSeqLong( + seqClass: CaseClassWithSeqLong, + setClass: CaseClassWithSetOfLongs) + +case class Blah(foo: String) + +case class TestIdStringWrapper(id: String) extends WrappedValue[String] + +case class ObjWithTestId(id: TestIdStringWrapper) + +object Obj { + + case class NestedCaseClassInObject(id: String) + + case class NestedCaseClassInObjectWithNestedCaseClassInObjectParam( + nested: NestedCaseClassInObject) + +} + +class TypeAndCompanion +object TypeAndCompanion { + case class NestedCaseClassInCompanion(id: String) +} + +case class WrapperValueClass(underlying: Int) extends AnyVal with WrappedValue[Int] + +case class WrappedValueInt(value: Int) extends WrappedValue[Int] + +case class WrappedValueLong(value: Long) extends WrappedValue[Long] + +case class WrappedValueString(value: String) extends WrappedValue[String] + +case class WrappedValueIntInObj(foo: WrappedValueInt) + +case class WrappedValueStringInObj(foo: WrappedValueString) + +case class WrappedValueLongInObj(foo: WrappedValueLong) + +case class CaseClassWithVal(name: String) { + val `type`: String = "person" +} + +case class CaseClassWithEnum(name: String, make: CarMakeEnum) + +case class CaseClassWithComplexEnums( + name: String, + make: CarMakeEnum, + makeOpt: Option[CarMakeEnum], + makeSeq: Seq[CarMakeEnum], + makeSet: Set[CarMakeEnum]) + +case class CaseClassWithSeqEnum(enumSeq: Seq[CarMakeEnum]) + +case class CaseClassWithOptionEnum(enumOpt: Option[CarMakeEnum]) + +case class CaseClassWithDateTime(dateTime: LocalDateTime) + +case class CaseClassWithIntAndDateTime( + @NotEmpty name: String, + age: Int, + age2: Int, + age3: Int, + dateTime: LocalDateTime, + dateTime2: LocalDateTime, + dateTime3: LocalDateTime, + dateTime4: LocalDateTime, + @NotEmpty dateTime5: Option[LocalDateTime]) + +case class CaseClassWithTwitterUtilDuration(duration: ctu.Duration) + +case class ClassWithFooClassInject(@JacksonInject fooClass: FooClass) + +case class ClassWithFooClassInjectAndDefault(@JacksonInject fooClass: FooClass = FooClass("12345")) + +case class ClassWithQueryParamDateTimeInject(@TestInjectableValue dateTime: Temporal) + +case class CaseClassWithEscapedLong(`1-5`: Long) + +case class CaseClassWithEscapedLongAndAnnotation(@Max(25) `1-5`: Long) + +case class CaseClassWithEscapedString(`1-5`: String) + +case class CaseClassWithEscapedNormalString(`a`: String) + +case class UnicodeNameCaseClass(`winning-id`: Int, name: String) + +case class TestEntityIdsResponse(entityIds: Seq[Long], previousCursor: String, nextCursor: String) + +object TestEntityIdsResponseWithCompanion { + val msg = "im the companion" +} + +case class TestEntityIdsResponseWithCompanion( + entityIds: Seq[Long], + previousCursor: String, + nextCursor: String) + +case class WrappedValueStringMapObject(map: Map[WrappedValueString, String]) + +case class FooClass(id: String) + +case class Group3(id: String) extends SerdeLogging + +case class CaseClassWithNotEmptyValidation(@NotEmpty name: String, make: CarMakeEnum) + +case class CaseClassWithInvalidValidation( + @ThrowsRuntimeExceptionConstraint name: String, + make: CarMakeEnum) + +// a constraint with no provided or registered validator class implementation +case class CaseClassWithNoValidatorImplConstraint( + @NoValidatorImplConstraint name: String, + make: CarMakeEnum) + +case class CaseClassWithMethodValidationException( + @NotEmpty name: String, + passphrase: String) { + @MethodValidation(fields = Array("passphrase")) + def checkPassphrase: MethodValidationResult = { + throw new RuntimeException("oh noes") + } +} + +case class NoConstructorArgs() + +case class CaseClassWithBoolean(foo: Boolean) + +case class CaseClassWithSeqBooleans(foos: Seq[Boolean]) + +case class CaseClassInjectStringWithDefault(@JacksonInject string: String = "DefaultHello") + +case class CaseClassInjectInt(@JacksonInject age: Int) + +case class CaseClassInjectOptionInt(@JacksonInject age: Option[Int]) + +case class CaseClassInjectOptionString(@JacksonInject string: Option[String]) + +case class CaseClassInjectString(@JacksonInject string: String) + +case class CaseClassWithCustomDecimalFormat( + @JsonDeserialize(using = classOf[MyBigDecimalDeserializer]) myBigDecimal: BigDecimal, + @JsonDeserialize(using = classOf[MyBigDecimalDeserializer]) optMyBigDecimal: Option[BigDecimal]) + +case class CaseClassWithLongAndDeserializer( + @JsonDeserialize(contentAs = classOf[java.lang.Long]) + long: Number) + +case class CaseClassWithOptionLongAndDeserializer( + @JsonDeserialize(contentAs = classOf[java.lang.Long]) + optLong: Option[Long]) + +case class CaseClassWithTwoConstructors(id: Long, @NotEmpty name: String) { + def this(id: Long) = this(id, "New User") +} + +case class CaseClassWithThreeConstructors(id: Long, @NotEmpty name: String) { + def this(id: Long) = this(id, "New User") + + def this(name: String) = this(42, name) +} + +case class Person( + id: Int, + @NotEmpty name: String, + dob: Option[LocalDate] = None, + age: Option[Int], + age_with_default: Option[Int] = None, + nickname: String = "unknown", + address: Option[Address] = None) + +class MyBigDecimalDeserializer extends StdDeserializer[BigDecimal](classOf[BigDecimal]) { + override def deserialize(jp: JsonParser, ctxt: DeserializationContext): BigDecimal = { + val jsonNode: ValueNode = jp.getCodec.readTree(jp) + BigDecimal(jsonNode.asText).setScale(2, RoundingMode.HALF_UP) + } + + override def getEmptyValue: BigDecimal = BigDecimal(0) +} + +// allows parsing a JSON null into a null type for this object +class NullableCarMakeDeserializer extends StdDeserializer[CarMakeEnum](classOf[CarMakeEnum]) { + override def deserialize(jp: JsonParser, ctxt: DeserializationContext): CarMakeEnum = { + val jsonNode: ValueNode = jp.getCodec.readTree(jp) + if (jsonNode.isNull) { + null + } else { + CarMakeEnum.valueOf(jsonNode.asText()) + } + } +} + +case class WithNullableCarMake( + nonNullCarMake: CarMakeEnum, + @JsonDeserialize(using = classOf[NullableCarMakeDeserializer]) nullableCarMake: CarMakeEnum) + +case class WithDefaultNullableCarMake( + @JsonDeserialize(using = classOf[NullableCarMakeDeserializer]) nullableCarMake: CarMakeEnum = + CarMakeEnum.ford) + +// is a string type but json types are passed as the value +case class WithJsonStringType(value: String) + +case class WithEmptyJsonProperty(@JsonProperty foo: String) + +case class WithNonemptyJsonProperty(@JsonProperty("bar") foo: String) + +case class WithoutJsonPropertyAnnotation(foo: String) + +case class NamingStrategyJsonProperty(@JsonProperty longFieldName: String) + +case class Address( + @NotEmpty street: Option[String] = None, + @NotEmpty city: String, + @NotEmpty state: String) { + + @MethodValidation + def validateState: MethodValidationResult = + MethodValidationResult.validIfTrue( + state == "CA" || state == "MD" || state == "WI", + "state must be one of [CA, MD, WI]" + ) +} + +trait CaseClassTrait { + @JsonProperty("fedoras") + @Size(min = 1, max = 2) + def names: Seq[String] + + @Min(1L) + def age: Int +} +case class CaseClassTraitImpl(names: Seq[String], @JsonProperty("oldness") age: Int) + extends CaseClassTrait + +package object internal { + case class SimplePersonInPackageObject( // not recommended but used here for testing use case + name: String = "default-name") + + case class SimplePersonInPackageObjectWithoutConstructorParams( + ) // not recommended but used here for testing use case +} + +case class LimiterProfile(id: Long, name: String) +case class LimiterProfiles(profiles: Option[Seq[LimiterProfile]] = None) + +object LimiterProfiles { + def apply(profiles: Seq[LimiterProfile]): LimiterProfiles = { + if (profiles.isEmpty) { + LimiterProfiles() + } else { + LimiterProfiles(Some(profiles)) + } + } +} + +case class CaseClassWithBoxedPrimitives(events: Integer, errors: Integer) + +sealed trait ClusterRequest +case class AddClusterRequest( + @Size(min = 0, max = 30) clusterName: String, + @NotEmpty job: String, + zone: String, + environment: String, + dtab: String, + address: String, + owners: String = "", + dedicated: Boolean = true, + enabled: Boolean = true, + description: String = "") + extends ClusterRequest { + + @MethodValidation(fields = Array("clusterName"), payload = Array(classOf[PatternNotMatched])) + def validateClusterName: MethodValidationResult = { + validateName(clusterName) + } + + private def validateName(name: String): MethodValidationResult = { + val regex = "[0-9a-zA-Z_\\-\\.>]+" + MethodValidationResult.validIfTrue( + name.matches(regex), + s"$name is invalid. Only alphanumeric and special characters from (_,-,.,>) are allowed.", + Some(PatternNotMatched(name, regex))) + } +} + +case class PatternNotMatched(pattern: String, regex: String) extends Payload + +case class Point(abscissa: Int, ordinate: Int) { + def area: Int = abscissa * ordinate +} + +case class Points(first: Point, second: Point) + +trait PointMixin { + @JsonProperty("x") @Min(0) @Max(100) def abscissa: Int + @JsonProperty("y") @Min(0) @Max(100) def ordinate: Int + @JsonIgnore def area: Int +} + +@JsonIgnoreType +case class IgnoreMe(id: Long) + +case class ContainsAnIgnoreTypeNoDefault(ignored: IgnoreMe, name: String, description: String) +case class ContainsAnIgnoreTypeWithDefault( + ignored: IgnoreMe = IgnoreMe(42L), + name: String, + description: String) + +object Views { + class Public + class Internal extends Public +} + +case class Item( + @JsonView(Array(classOf[Views.Public])) id: Long = 1L, + @JsonView(Array(classOf[Views.Public])) name: String = "", + @JsonView(Array(classOf[Views.Internal])) owner: String = "") + +case class ItemSomeViews( + @JsonView(Array(classOf[Views.Public])) id: Long = 1L, + @JsonView(Array(classOf[Views.Public])) name: String = "", + owner: String = "") + +case class ItemNoDefaultForView(@JsonView(Array(classOf[Views.Public])) name: String) + +@JsonNaming(classOf[PropertyNamingStrategies.KebabCaseStrategy]) +case class CaseClassWithKebabCase(pleaseUseKebabCase: Boolean) + +@JsonNaming(classOf[PropertyNamingStrategies.KebabCaseStrategy]) +trait KebabCaseMixin + +case class CaseClassShouldUseKebabCaseFromMixin(willThisGetTheRightCasing: Boolean) + +@JsonNaming +case class UseDefaultNamingStrategy(thisFieldShouldUseDefaultPropertyNamingStrategy: Boolean) diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/serde/DurationStringDeserializerTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/DurationStringDeserializerTest.scala new file mode 100644 index 0000000000..2b36c40ade --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/DurationStringDeserializerTest.scala @@ -0,0 +1,52 @@ +package com.twitter.util.jackson.serde + +import com.twitter.conversions.DurationOps._ +import com.twitter.util.Duration +import com.twitter.util.jackson.ScalaObjectMapper +import com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +case class CaseClassWithDuration(d: Duration) + +@RunWith(classOf[JUnitRunner]) +class DurationStringDeserializerTest extends AnyFunSuite with Matchers { + + private[this] val mapper = ScalaObjectMapper() + + test("deserializes values") { + Seq( + " 1.second" -> 1.second, + "+1.second" -> 1.second, + "-1.second" -> -1.second, + "1.SECOND" -> 1.second, + "1.day - 1.second" -> (1.day - 1.second), + "1.day" -> 1.day, + "1.microsecond" -> 1.microsecond, + "1.millisecond" -> 1.millisecond, + "1.second" -> 1.second, + "1.second+1.minute + 1.day" -> (1.second + 1.minute + 1.day), + "1.second+1.second" -> 2.seconds, + "2.hours" -> 2.hours, + "3.days" -> 3.days, + "321.nanoseconds" -> 321.nanoseconds, + "65.minutes" -> 65.minutes, + "876.milliseconds" -> 876.milliseconds, + "98.seconds" -> 98.seconds, + "Duration.Bottom" -> Duration.Bottom, + "Duration.Top" -> Duration.Top, + "Duration.Undefined" -> Duration.Undefined, + "duration.TOP" -> Duration.Top + ) foreach { + case (s, d) => + mapper.parse[CaseClassWithDuration](s"""{"d":"$s"}""") should equal( + CaseClassWithDuration(d)) + } + + intercept[CaseClassMappingException] { + mapper.parse[CaseClassWithDuration]("{}") + } + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/serde/LocalDateTimeSerdeTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/LocalDateTimeSerdeTest.scala new file mode 100644 index 0000000000..c972f504db --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/LocalDateTimeSerdeTest.scala @@ -0,0 +1,125 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.databind.{SerializationFeature, ObjectMapper => JacksonObjectMapper} +import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule +import com.fasterxml.jackson.module.scala.{ + DefaultScalaModule, + ScalaObjectMapper => JacksonScalaObjectMapper +} +import com.twitter.util.jackson.ScalaObjectMapper +import java.time.format.DateTimeFormatter +import java.time.{LocalDateTime, ZoneId} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +case class MyClass(dateTime: LocalDateTime) + +private object LocalDateTimeSerdeTest { + val UTCZoneId: ZoneId = ZoneId.of("UTC") + val ISO8601DateTimeFormatter: DateTimeFormatter = + DateTimeFormatter.ISO_OFFSET_DATE_TIME // "2014-05-30T03:57:59.302Z" + + // 1st one for the LocalDate, the 2nd one for LocalDateTime + // @JsonFormat(pattern = "yyyy-MM-dd") + // @JsonFormat(pattern = "yyyy-MM-dd HH:mm") +} + +@RunWith(classOf[JUnitRunner]) +class LocalDateTimeSerdeTest extends AnyFunSuite with Matchers { + import LocalDateTimeSerdeTest._ + + private final val NowUtc: LocalDateTime = LocalDateTime.now(UTCZoneId) + private[this] val myClass1 = MyClass(dateTime = NowUtc) + + // Most {@code java.time} types are serialized as numbers (integers or decimals as appropriate) if the + // {@link com.fasterxml.jackson.databind.SerializationFeature#WRITE_DATES_AS_TIMESTAMPS} feature is enabled, and otherwise are serialized in + // standard ISO-8601 string representation. ISO-8601 specifies formats + // for representing offset dates and times, zoned dates and times, local dates and times, periods, durations, zones, and more. All + // {@code java.time} types have built-in translation to and from ISO-8601 formats. + + // Granularity of timestamps is controlled through the companion features + // {@link com.fasterxml.jackson.databind.SerializationFeature#WRITE_DATE_TIMESTAMPS_AS_NANOSECONDS} and + // {@link com.fasterxml.jackson.databind.DeserializationFeature#READ_DATE_TIMESTAMPS_AS_NANOSECONDS}. For serialization, timestamps are + // written as fractional numbers (decimals), where the number is seconds and the decimal is fractional seconds, if + // {@code WRITE_DATE_TIMESTAMPS_AS_NANOSECONDS} is enabled (it is by default), with resolution as fine as nanoseconds depending on the + // underlying JDK implementation. If {@code WRITE_DATE_TIMESTAMPS_AS_NANOSECONDS} is disabled, timestamps are written as a whole number of + // milliseconds. At deserialization time, decimal numbers are always read as fractional second timestamps with up-to-nanosecond resolution, + // since the meaning of the decimal is unambiguous. The more ambiguous integer types are read as fractional seconds without a decimal point + // if {@code READ_DATE_TIMESTAMPS_AS_NANOSECONDS} is enabled (it is by default), and otherwise they are read as milliseconds. + + // {@link LocalDate}, {@link LocalTime}, {@link LocalDateTime}, and {@link OffsetTime}, which cannot portably be converted to timestamps + // and are instead represented as arrays when WRITE_DATES_AS_TIMESTAMPS is enabled. + // https://github.com/FasterXML/jackson-datatype-jsr310/blob/master/src/main/java/com/fasterxml/jackson/datatype/jsr310/JavaTimeModule.java + test("serialize as array") { + val mapper = new JacksonObjectMapper() + mapper.configure(SerializationFeature.WRITE_DATES_AS_TIMESTAMPS, true) + mapper.registerModule(new JavaTimeModule) + val serialized = mapper.writeValueAsString(NowUtc) // [2021,5,10,0,0,0,0] + // remove beginning [ and ending ] from the string + val serializedParts = serialized.substring(1, serialized.length - 1).split(",") + serializedParts.size should equal(7) + serializedParts(0) should equal(NowUtc.getYear.toString) + serializedParts(1) should equal(NowUtc.getMonthValue.toString) + serializedParts(2) should equal(NowUtc.getDayOfMonth.toString) + serializedParts(3) should equal(NowUtc.getHour.toString) + serializedParts(4) should equal(NowUtc.getMinute.toString) + serializedParts(5) should equal(NowUtc.getSecond.toString) + serializedParts(6) should equal(NowUtc.getNano.toString) + } + + test("serialize as text") { + val mapper = new JacksonObjectMapper() + mapper.configure(SerializationFeature.WRITE_DATES_AS_TIMESTAMPS, false) + mapper.registerModule(new JavaTimeModule) + val actualStr = mapper.writeValueAsString(NowUtc) + val actualStrUnquoted = + actualStr.substring(1, actualStr.length - 1) // drop beginning and end quotes + val actual = + actualStrUnquoted.substring( + 0, + actualStrUnquoted.lastIndexOf('.') + ) // // drop fractions of second + val expectStr = NowUtc.toString + val expected = expectStr.substring(0, expectStr.lastIndexOf('.')) // drop fractions of second + expected should be(actual) + } + + test("deserialize from text") { + val mapper = new JacksonObjectMapper with JacksonScalaObjectMapper + mapper.registerModule(new JavaTimeModule) + mapper.registerModule(DefaultSerdeModule) + mapper.readValue[LocalDateTime](quote(NowUtc.toString)) should equal(NowUtc) + } + + test("deserialize from text with ScalaObjectMapper with SerdeModule") { + val mapper = ScalaObjectMapper() + mapper.parse[LocalDateTime](quote(NowUtc.toString)) should equal(NowUtc) + } + + test("roundtrip text") { + val mapper = new JacksonObjectMapper with JacksonScalaObjectMapper + mapper.registerModule(DefaultScalaModule) + mapper.registerModule(new JavaTimeModule) + val json = mapper.writeValueAsString(myClass1) + val deserializedC1 = mapper.readValue[MyClass](json) + deserializedC1 should equal(myClass1) + } + + test("roundtrip text with ScalaObjectMapper with SerdeModule") { + val mapper = ScalaObjectMapper() + val json = mapper.writeValueAsString(myClass1) + val deserializedC1 = mapper.parse[MyClass](json) + deserializedC1 should equal(myClass1) + } + + test("roundtrip text with ScalaObjectMapper") { + val mapper = ScalaObjectMapper() + val json = mapper.writeValueAsString(myClass1) + val deserializedC1 = mapper.parse[MyClass](json) + deserializedC1 should equal(myClass1) + } + + private[this] def quote(str: String): String = s"""\"$str\"""" +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringDeserializerTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringDeserializerTest.scala new file mode 100644 index 0000000000..bfd04f9329 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringDeserializerTest.scala @@ -0,0 +1,113 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.annotation.JsonFormat +import com.fasterxml.jackson.databind.ObjectMapper +import com.fasterxml.jackson.module.scala.DefaultScalaModule +import com.twitter.util.jackson.ScalaObjectMapper +import com.twitter.util.jackson.caseclass.exceptions.CaseClassMappingException +import com.twitter.util.{Time, TimeFormat} +import org.junit.runner.RunWith +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +private object TimeStringDeserializerTest { + case class WithoutJsonFormat(time: Time) + case class WithJsonFormat(@JsonFormat(pattern = "yyyy-MM-dd") time: Time) + case class WithJsonFormatAndTimezone( + @JsonFormat(pattern = "yyyy-MM-dd HH:mm:ss", timezone = "America/Los_Angeles") time: Time) + case class WithEpochJsonFormat( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "s") time: Time) +} + +@RunWith(classOf[JUnitRunner]) +class TimeStringDeserializerTest extends AnyFunSuite with BeforeAndAfterAll with Matchers { + import TimeStringDeserializerTest._ + + private[this] final val Input1 = + """ + |{ + | "time": "2019-06-17T15:45:00.000+0000" + |} + """.stripMargin + + private[this] final val Input2 = + """ + |{ + | "time": "2019-06-17" + |} + """.stripMargin + + private[this] final val Input3 = + """ + |{ + | "time": "2019-06-17 16:30:00" + |} + """.stripMargin + + private[this] final val EpochInput = + """ + |{ + | "time": 1560786300 + |} + """.stripMargin + + private[this] val jacksonObjectMapper = new ObjectMapper() + private[this] val objectMapper = ScalaObjectMapper() + + override def beforeAll(): Unit = { + val modules = Seq(DefaultScalaModule, DefaultSerdeModule) + modules.foreach(jacksonObjectMapper.registerModule) + } + + test("should deserialize date without JsonFormat") { + val expected = 1560786300000L + val jacksonActual: WithoutJsonFormat = + jacksonObjectMapper.readValue(Input1, classOf[WithoutJsonFormat]) + jacksonActual.time.inMillis shouldEqual expected + + val frameworkActual = objectMapper.parse[WithoutJsonFormat](Input1) + frameworkActual.time.inMillis shouldEqual expected + } + + test("should deserialize date with JsonFormat") { + val expected: Time = new TimeFormat("yyyy-MM-dd").parse("2019-06-17") + val jacksonActual: WithJsonFormat = + jacksonObjectMapper.readValue(Input2, classOf[WithJsonFormat]) + jacksonActual.time shouldEqual expected + + val frameworkActual = objectMapper.parse[WithJsonFormat](Input2) + frameworkActual.time shouldEqual expected + } + + test("should deserialize date with JsonFormat and timezone") { + val expected = 1560814200000L + val jacksonActual: WithJsonFormatAndTimezone = + jacksonObjectMapper.readValue(Input3, classOf[WithJsonFormatAndTimezone]) + jacksonActual.time.inMillis shouldEqual expected + + val frameworkActual = objectMapper.parse[WithJsonFormatAndTimezone](Input3) + frameworkActual.time.inMillis shouldEqual expected + } + + test("should return a MappingException for empty value") { + val jacksonActual: WithoutJsonFormat = + jacksonObjectMapper.readValue("{}", classOf[WithoutJsonFormat]) + jacksonActual should equal(WithoutJsonFormat(null)) // underlying mapper allows value to be null + + intercept[CaseClassMappingException] { + objectMapper.parse[WithoutJsonFormat]("{}") + } + } + + test("should support JsonFormat with epoch timestamp") { + val expected = 1560786300000L + val jacksonActual: WithEpochJsonFormat = + jacksonObjectMapper.readValue(EpochInput, classOf[WithEpochJsonFormat]) + jacksonActual.time.inMillis shouldEqual expected + + val frameworkActual = objectMapper.parse[WithEpochJsonFormat](EpochInput) + frameworkActual.time.inMillis shouldEqual expected + } +} diff --git a/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringSerializerTest.scala b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringSerializerTest.scala new file mode 100644 index 0000000000..2ae594cd14 --- /dev/null +++ b/util-jackson/src/test/scala/com/twitter/util/jackson/serde/TimeStringSerializerTest.scala @@ -0,0 +1,139 @@ +package com.twitter.util.jackson.serde + +import com.fasterxml.jackson.annotation.JsonFormat +import com.fasterxml.jackson.databind.ObjectMapper +import com.fasterxml.jackson.module.scala.DefaultScalaModule +import com.twitter.util.jackson.ScalaObjectMapper +import com.twitter.util.{Time, TimeFormat} +import java.util.TimeZone +import org.junit.runner.RunWith +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +private object TimeStringSerializerTest { + case class WithoutJsonFormat(time: Time) + case class WithJsonFormat(@JsonFormat(pattern = "yyyy-MM-dd") time: Time) + // NOTE: JsonFormat.Shape is not respected, times are always written as Strings + case class WithInvalidJsonFormat( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "yyyy-MM-dd") time: Time) + case class WithJsonFormatAndTimezone( + @JsonFormat(pattern = "yyyy-MM-dd HH:mm:ss", timezone = "America/Los_Angeles") time: Time) + // Note for serialization the pattern 's' represents the "second in minute": see: https://docs.oracle.com/javase/7/docs/api/java/text/SimpleDateFormat.html + // s Second in minute Number 55 + case class WithEpochJsonFormat( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "s") time: Time) + // S Millisecond Number 978 + case class WithJsonFormatNumberShape( + @JsonFormat(shape = JsonFormat.Shape.NUMBER_INT, pattern = "S") time: Time) + case class WithJsonFormatNaturalShape( + @JsonFormat(shape = JsonFormat.Shape.NATURAL, pattern = "S") time: Time) + case class WithJsonFormatNoShape( + @JsonFormat(shape = JsonFormat.Shape.ANY, pattern = "S") time: Time) + case class JsonFormatNumberShape( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "s.SS") time: Time) + case class JsonFormatNumberShape2( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "yyyy.MM.dd") time: Time) + case class JsonFormatNumberShape3( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "dd") time: Time) + case class JsonFormatNumberShape4( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "EEE") time: Time) + case class JsonFormatNumberShape5( + @JsonFormat(shape = JsonFormat.Shape.NUMBER, pattern = "yyyy,MM,dd") time: Time) + + val DefaultTimeFormat: TimeFormat = new TimeFormat("s") + val DefaultTime: Time = DefaultTimeFormat.parse("1560786300") // "2019-06-17 15:45:00 +0000" +} + +@RunWith(classOf[JUnitRunner]) +class TimeStringSerializerTest extends AnyFunSuite with BeforeAndAfterAll with Matchers { + import TimeStringSerializerTest._ + + private[this] val jacksonObjectMapper = new ObjectMapper() + private[this] val objectMapper = ScalaObjectMapper() + + override def beforeAll(): Unit = { + val modules = Seq(DefaultScalaModule, DefaultSerdeModule) + modules.foreach(jacksonObjectMapper.registerModule) + } + + test("should serialize date without JsonFormat") { + val expected: String = + """{"time":"2019-06-17 15:45:00 +0000"}""" + + val value = WithoutJsonFormat(DefaultTime) + jacksonObjectMapper.writeValueAsString(value) shouldEqual expected + objectMapper.writeValueAsString(value) shouldEqual expected + } + + test("should serialize as String with invalid JsonFormat") { + val expected: String = + """{"time":"2019-06-17"}""" + + val value = WithInvalidJsonFormat(DefaultTime) + jacksonObjectMapper.writeValueAsString(value) shouldEqual expected + objectMapper.writeValueAsString(value) shouldEqual expected + } + + test("should serialize date with JsonFormat and timezone") { + val expected = """{"time":"2019-06-17 23:30:00"}""" + + val time = + new TimeFormat( + pattern = "s", + locale = None, + timezone = TimeZone.getTimeZone("America/Los_Angeles")).parse("1560814200") + val value = WithJsonFormatAndTimezone(time) + jacksonObjectMapper.writeValueAsString(value) shouldEqual expected + objectMapper.writeValueAsString(value) shouldEqual expected + } + + test("should serialize date with JsonFormat") { + val value0 = WithJsonFormat(DefaultTime) + jacksonObjectMapper.writeValueAsString(value0) shouldEqual """{"time":"2019-06-17"}""" + objectMapper.writeValueAsString(value0) shouldEqual """{"time":"2019-06-17"}""" + + // Note the @JsonFormat only ends up formatting the + // "second in minute", see: https://docs.oracle.com/javase/7/docs/api/java/text/SimpleDateFormat.html + // which in this example is 0. + val value = WithEpochJsonFormat(DefaultTime) + jacksonObjectMapper.writeValueAsString(value) shouldEqual """{"time":"0"}""" + objectMapper.writeValueAsString(value) shouldEqual """{"time":"0"}""" + + // Note the @JsonFormat only ends up formatting the + // "millisecond", see: https://docs.oracle.com/javase/7/docs/api/java/text/SimpleDateFormat.html + // which in this example is 0. + val value1 = WithJsonFormatNumberShape(DefaultTime) + jacksonObjectMapper.writeValueAsString(value1) shouldEqual """{"time":"0"}""" + objectMapper.writeValueAsString(value1) shouldEqual """{"time":"0"}""" + + val value2 = WithJsonFormatNaturalShape(DefaultTime) + jacksonObjectMapper.writeValueAsString(value2) shouldEqual """{"time":"0"}""" + objectMapper.writeValueAsString(value2) shouldEqual """{"time":"0"}""" + + val value3 = WithJsonFormatNoShape(DefaultTime) + jacksonObjectMapper.writeValueAsString(value3) shouldEqual """{"time":"0"}""" + objectMapper.writeValueAsString(value3) shouldEqual """{"time":"0"}""" + + val value4 = JsonFormatNumberShape(DefaultTime) + jacksonObjectMapper.writeValueAsString(value4) shouldEqual """{"time":"0.00"}""" + objectMapper.writeValueAsString(value4) shouldEqual """{"time":"0.00"}""" + + val value5 = JsonFormatNumberShape2(DefaultTime) + jacksonObjectMapper.writeValueAsString(value5) shouldEqual """{"time":"2019.06.17"}""" + objectMapper.writeValueAsString(value5) shouldEqual """{"time":"2019.06.17"}""" + + val value6 = JsonFormatNumberShape3(DefaultTime) + jacksonObjectMapper.writeValueAsString(value6) shouldEqual """{"time":"17"}""" + objectMapper.writeValueAsString(value6) shouldEqual """{"time":"17"}""" + + val value7 = JsonFormatNumberShape4(DefaultTime) + jacksonObjectMapper.writeValueAsString(value7) shouldEqual """{"time":"Mon"}""" + objectMapper.writeValueAsString(value7) shouldEqual """{"time":"Mon"}""" + + val value8 = JsonFormatNumberShape5(DefaultTime) + jacksonObjectMapper.writeValueAsString(value8) shouldEqual """{"time":"2019,06,17"}""" + objectMapper.writeValueAsString(value8) shouldEqual """{"time":"2019,06,17"}""" + } +} diff --git a/util-jvm/BUILD b/util-jvm/BUILD index 5b41ae4712..f756c4066d 100644 --- a/util-jvm/BUILD +++ b/util-jvm/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-jvm/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-jvm/src/test/java", "util/util-jvm/src/test/scala", diff --git a/util-jvm/src/main/scala/BUILD b/util-jvm/src/main/scala/BUILD index 2d1e2415fa..0f5e600d28 100644 --- a/util-jvm/src/main/scala/BUILD +++ b/util-jvm/src/main/scala/BUILD @@ -1,17 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-jvm", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-stats/src/main/scala", - ], - exports = [ - "util/util-core/src/main/scala", + "util/util-jvm/src/main/scala/com/twitter/jvm", ], ) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Allocations.scala b/util-jvm/src/main/scala/com/twitter/jvm/Allocations.scala index b33089dee1..996a65baae 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/Allocations.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/Allocations.scala @@ -1,19 +1,20 @@ package com.twitter.jvm import com.twitter.finagle.stats.StatsReceiver -import java.lang.management.{MemoryPoolMXBean, MemoryUsage, ManagementFactory} +import java.lang.management.MemoryPoolMXBean +import java.lang.management.MemoryUsage +import java.lang.management.ManagementFactory import java.util.concurrent.LinkedBlockingQueue import java.util.concurrent.atomic.AtomicLong import java.util.{List => juList} -import javax.management.openmbean.{CompositeData, TabularData} -import javax.management.{ - ListenerNotFoundException, - Notification, - NotificationListener, - NotificationEmitter -} -import scala.collection.JavaConverters._ +import javax.management.openmbean.CompositeData +import javax.management.openmbean.TabularData +import javax.management.ListenerNotFoundException +import javax.management.Notification +import javax.management.NotificationListener +import javax.management.NotificationEmitter import scala.collection.mutable +import scala.jdk.CollectionConverters._ private[jvm] object Allocations { val Unknown: Long = -1L @@ -30,8 +31,7 @@ private[jvm] class Allocations(statsReceiver: StatsReceiver) { private[this] val edenPool: Option[MemoryPoolMXBean] = ManagementFactory.getMemoryPoolMXBeans.asScala.find { bean => - // todo: see if we can support the g1 collector - bean.getName == "Par Eden Space" || bean.getName == "PS Eden Space" + bean.getName == "Par Eden Space" || bean.getName == "PS Eden Space" || bean.getName == "G1 Eden Space" } private[this] val edenSizeAfterLastGc = new AtomicLong() @@ -41,13 +41,12 @@ private[jvm] class Allocations(statsReceiver: StatsReceiver) { private[this] val beanAndListeners = new LinkedBlockingQueue[(NotificationEmitter, NotificationListener)]() - private[this] val edenGcPauses = statsReceiver.scope("eden").stat("pause_msec") + private[this] val edenGcPauses = statsReceiver + .hierarchicalScope("eden").label("jmm", "eden").stat("pause_msec") private[jvm] def start(): Unit = { edenPool - .flatMap { bean => - Option(bean.getUsage) - } + .flatMap { bean => Option(bean.getUsage) } .foreach { _ => ManagementFactory.getGarbageCollectorMXBeans.asScala.foreach { case bean: NotificationEmitter => diff --git a/util-jvm/src/main/scala/com/twitter/jvm/BUILD b/util-jvm/src/main/scala/com/twitter/jvm/BUILD new file mode 100644 index 0000000000..abdf640814 --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/BUILD @@ -0,0 +1,20 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-jvm", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-app/src/main/scala", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", + ], + exports = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/ContentionSnapshot.scala b/util-jvm/src/main/scala/com/twitter/jvm/ContentionSnapshot.scala index fe2de5980a..8e189a0ed1 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/ContentionSnapshot.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/ContentionSnapshot.scala @@ -5,10 +5,15 @@ import java.lang.management.{ManagementFactory, ThreadInfo} /** * A thread contention summary providing a brief overview of threads - * that are [[http://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#BLOCKED BLOCKED]], [[http://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#WAITING WAITING]], or [[http://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#TIMED_WAITING TIMED_WAITING]] + * that are [[https://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#BLOCKED BLOCKED]], + * [[https://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#WAITING WAITING]], + * or [[https://docs.oracle.com/javase/1.5.0/docs/api/java/lang/Thread.State.html#TIMED_WAITING TIMED_WAITING]] * * While this could be an object, we use instantiation as a signal of * intent and enable contention monitoring. + * + * @note users should ensure that the `java.lang.management.ManagementPermission("control")` is + * allowed via the `java.lang.SecurityManager`. */ class ContentionSnapshot { ManagementFactory.getThreadMXBean.setThreadContentionMonitoringEnabled(true) @@ -43,9 +48,7 @@ class ContentionSnapshot { if (deadlockThreadIds == null) Array.empty[ThreadInfo] else deadlockThreadIds.flatMap { id => - blocked.find { threadInfo => - threadInfo.getThreadId() == id - } + blocked.find { threadInfo => threadInfo.getThreadId() == id } } Snapshot( diff --git a/util-jvm/src/main/scala/com/twitter/jvm/CpuProfile.scala b/util-jvm/src/main/scala/com/twitter/jvm/CpuProfile.scala index 02b802094f..fff58c38ac 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/CpuProfile.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/CpuProfile.scala @@ -1,14 +1,12 @@ package com.twitter.jvm +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Duration, Future, Promise, Stopwatch, Time} import java.io.OutputStream import java.lang.management.ManagementFactory import java.nio.{ByteBuffer, ByteOrder} - import scala.collection.mutable -import com.twitter.conversions.time._ -import com.twitter.util.{Duration, Future, Promise, Stopwatch, Time} - /** * A CPU profile. */ @@ -20,14 +18,13 @@ case class CpuProfile( // The number of samples taken. count: Int, // The number of samples missed. - missed: Int -) { + missed: Int) { /** * Write a Google pprof-compatible profile to `out`. The format is * documented here: * - * http://google-perftools.googlecode.com/svn/trunk/doc/cpuprofile-fileformat.html + * https://google-perftools.googlecode.com/svn/trunk/doc/cpuprofile-fileformat.html */ def writeGoogleProfile(out: OutputStream): Unit = { var next = 1 @@ -83,7 +80,7 @@ object CpuProfile { * When looking for RUNNABLEs, the JVM's notion of runnable differs from the * from kernel's definition and for some well known cases, we can filter * out threads that are actually asleep. - * See http://www.brendangregg.com/blog/2014-06-09/java-cpu-sampling-using-hprof.html + * See https://www.brendangregg.com/blog/2014-06-09/java-cpu-sampling-using-hprof.html */ private[jvm] def isRunnable(stackElem: StackTraceElement): Boolean = !IdleClassAndMethod.contains((stackElem.getClassName, stackElem.getMethodName)) @@ -119,7 +116,7 @@ object CpuProfile { val bean = ManagementFactory.getThreadMXBean() val elapsed = Stopwatch.start() val end = howlong.fromNow - val period = (1000000 / frequency).microseconds + val period = ((1000000L / frequency)).microseconds val myId = Thread.currentThread().getId() var next = Time.now diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Estimator.scala b/util-jvm/src/main/scala/com/twitter/jvm/Estimator.scala index 2e64022f67..c7e078f835 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/Estimator.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/Estimator.scala @@ -40,7 +40,7 @@ class Kalman(N: Int) { val mv = mvar val ev = evar if (mv + ev == 0) - weight = 1D + weight = 1d else weight = mv / (mv + ev) n += 1 @@ -56,21 +56,19 @@ class Kalman(N: Int) { else ebuf ) - def estimate = est + def estimate: Double = est private[this] def variance(samples: Array[Double]): Double = { if (samples.length == 1) - return 0D + return 0d val sum = samples.sum val mean = sum / samples.length - val diff = (samples map { x => - (x - mean) * (x - mean) - }).sum + val diff = (samples map { x => (x - mean) * (x - mean) }).sum diff / (samples.length - 1) } - override def toString = + override def toString: String = "Kalman".format(estimate, weight, mvar, evar) } @@ -80,7 +78,7 @@ class Kalman(N: Int) { * value). */ class KalmanGaussianError(N: Int, range: Double) extends Kalman(N) with Estimator[Double] { - require(range >= 0D && range < 1D) + require(range >= 0d && range < 1d) private[this] val rng = new Random def measure(m: Double): Unit = { @@ -123,7 +121,7 @@ class WindowedMeans(N: Int, windows: Seq[(Int, Int)]) extends Estimator[Double] n += 1 } - def estimate = { + def estimate: Double = { require(n > 0) val weightedMeans = normalized map { case (w, i) => w * mean(n, i) } weightedMeans.sum @@ -134,10 +132,10 @@ class WindowedMeans(N: Int, windows: Seq[(Int, Int)]) extends Estimator[Double] * Unix-like load average, an exponentially weighted moving average, * smoothed to the given interval (counted in number of * measurements). - * See: http://web.mit.edu/saltzer/www/publications/instrumentation.html + * See: https://web.mit.edu/saltzer/www/publications/instrumentation.html */ class LoadAverage(interval: Double) extends Estimator[Double] { - private[this] val a = math.exp(-1D / interval) + private[this] val a = math.exp(-1d / interval) private[this] var load = Double.NaN def measure(m: Double): Unit = { @@ -146,5 +144,5 @@ class LoadAverage(interval: Double) extends Estimator[Double] { else load * a + m * (1 - a) } - def estimate = load + def estimate: Double = load } diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Gc.scala b/util-jvm/src/main/scala/com/twitter/jvm/Gc.scala new file mode 100644 index 0000000000..bf96b655bd --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/Gc.scala @@ -0,0 +1,5 @@ +package com.twitter.jvm + +import com.twitter.util.{Duration, Time} + +case class Gc(count: Long, name: String, timestamp: Time, duration: Duration) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/GcPredictor.scala b/util-jvm/src/main/scala/com/twitter/jvm/GcPredictor.scala index 477789003a..0b6d9b0c4f 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/GcPredictor.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/GcPredictor.scala @@ -1,6 +1,6 @@ package com.twitter.jvm -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Duration, Time, Timer} /** diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Heap.scala b/util-jvm/src/main/scala/com/twitter/jvm/Heap.scala new file mode 100644 index 0000000000..f75256695d --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/Heap.scala @@ -0,0 +1,15 @@ +package com.twitter.jvm + +/** + * Information about the Heap + * + * @param allocated Estimated number of bytes + * that have been allocated so far (into eden) + * + * @param tenuringThreshold How many times an Object + * needs to be copied before being tenured. + * + * @param ageHisto Histogram of the number of bytes that + * have been copied as many times. Note: 0-indexed. + */ +case class Heap(allocated: Long, tenuringThreshold: Long, ageHisto: Seq[Long]) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Heapster.scala b/util-jvm/src/main/scala/com/twitter/jvm/Heapster.scala index d39e38ccdf..07f9f6f910 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/Heapster.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/Heapster.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -34,12 +34,18 @@ class Heapster(klass: Class[_]) { def start(): Unit = { startM.invoke(null) } def shutdown(): Unit = { stopM.invoke(null) } - def setSamplingPeriod(period: java.lang.Integer): Unit = { setSamplingPeriodM.invoke(null, period) } + def setSamplingPeriod(period: java.lang.Integer): Unit = { + setSamplingPeriodM.invoke(null, period) + } def clearProfile(): Unit = { clearProfileM.invoke(null) } def dumpProfile(forceGC: java.lang.Boolean): Array[Byte] = dumpProfileM.invoke(null, forceGC).asInstanceOf[Array[Byte]] - def profile(howlong: Duration, samplingPeriod: Int = 10 << 19, forceGC: Boolean = true) = { + def profile( + howlong: Duration, + samplingPeriod: Int = 10 << 19, + forceGC: Boolean = true + ): Array[Byte] = { clearProfile() setSamplingPeriod(samplingPeriod) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Hotspot.scala b/util-jvm/src/main/scala/com/twitter/jvm/Hotspot.scala index e4b91e66d8..9b9c50e4df 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/Hotspot.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/Hotspot.scala @@ -1,18 +1,25 @@ package com.twitter.jvm -import com.twitter.conversions.storage._ -import com.twitter.conversions.time._ +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.Time import java.lang.management.ManagementFactory -import java.util.logging.{Level, Logger} +import java.util.logging.Level +import java.util.logging.Logger +import javax.management.openmbean.CompositeData import javax.management.openmbean.CompositeDataSupport -import javax.management.{ObjectName, RuntimeMBeanException} -import scala.collection.JavaConverters._ +import javax.management.Notification +import javax.management.NotificationEmitter +import javax.management.NotificationListener +import javax.management.ObjectName +import javax.management.RuntimeMBeanException +import javax.naming.OperationNotSupportedException +import scala.jdk.CollectionConverters._ import scala.language.reflectiveCalls class Hotspot extends Jvm { private[this] val epoch = - Time.fromMilliseconds(ManagementFactory.getRuntimeMXBean().getStartTime()) + Time.fromMilliseconds(ManagementFactory.getRuntimeMXBean.getStartTime) private[this] type Counter = { def getName(): String @@ -28,22 +35,16 @@ class Hotspot extends Jvm { ObjectName.getInstance("com.sun.management:type=HotSpotDiagnostic") private[this] val jvm: VMManagement = { - val fld = try { - // jdk5/6 have jvm field in ManagementFactory class - Class.forName("sun.management.ManagementFactory").getDeclaredField("jvm") - } catch { - case _: NoSuchFieldException => - // jdk7 moves jvm field to ManagementFactoryHelper class - Class.forName("sun.management.ManagementFactoryHelper").getDeclaredField("jvm") - } + val fld = Class + .forName("sun.management.ManagementFactoryHelper") + .getDeclaredField("jvm") fld.setAccessible(true) fld.get(null).asInstanceOf[VMManagement] } private[this] def opt(name: String) = try Some { - val o = ManagementFactory - .getPlatformMBeanServer() + val o = ManagementFactory.getPlatformMBeanServer .invoke(DiagnosticBean, "getVMOption", Array(name), Array("java.lang.String")) o.asInstanceOf[CompositeDataSupport].get("value").asInstanceOf[String] } catch { @@ -56,21 +57,29 @@ class Hotspot extends Jvm { private[this] def long(c: Counter) = c.getValue().asInstanceOf[Long] private[this] def counters(pat: String) = { - val cs = jvm.getInternalCounters(pat).asScala - cs.map { c => - c.getName() -> c - }.toMap + try { + val cs = jvm.getInternalCounters(pat).asScala + cs.map { c => c.getName() -> c }.toMap + } catch { + case e: OperationNotSupportedException => + log.log( + Level.WARNING, + s"failed to get internal JVM counters as these are not available on your JVM.") + Map.empty[String, Counter] + case e: Throwable => + throw e + } } private[this] def counter(name: String): Option[Counter] = counters(name).get(name) object opts extends Opts { - def compileThresh = opt("CompileThreshold") map (_.toInt) + def compileThresh: Option[Int] = opt("CompileThreshold").map(_.toInt) } private[this] def ticksToDuration(ticks: Long, freq: Long) = - (1000000 * ticks / freq).microseconds + (1000000L * ticks / freq).microseconds private[this] def getGc(which: Int, cs: Map[String, Counter]) = { def get(what: String) = cs.get("sun.gc.collector.%d.%s".format(which, what)) @@ -87,7 +96,26 @@ class Hotspot extends Jvm { } yield Gc(invocations, kind, epoch + lastEntryTime, duration) } - def snap: Snapshot = { + private[this] def newGcListener(f: Gc => Unit): NotificationListener = + new NotificationListener { + def handleNotification(notification: Notification, handback: Any): Unit = { + if (Hotspot.isGcNotification(notification)) { + val gc = + Hotspot.gcFromNotificationInfo(notification.getUserData.asInstanceOf[CompositeData]) + f(gc) + } + } + } + + override def foreachGc(f: Gc => Unit): Unit = { + ManagementFactory.getGarbageCollectorMXBeans.asScala.foreach { + case bean: NotificationEmitter => + bean.addNotificationListener(newGcListener(f), null, null) + case _ => () + } + } + + def snap: Snapshot = synchronized { val cs = counters("") val heap = for { invocations <- cs.get("sun.gc.collector.0.invocations").map(long) @@ -102,7 +130,7 @@ class Hotspot extends Jvm { val ageHisto = for { thresh <- tenuringThreshold.toSeq - i <- 1L to thresh + i <- 1L.to(thresh) bucket <- cs.get("sun.gc.generation.0.agetable.bytes.%02d".format(i)) } yield long(bucket) @@ -117,7 +145,7 @@ class Hotspot extends Jvm { // TODO: include causes for GCs? Snapshot( timestamp.getOrElse(Time.epoch), - heap.getOrElse(Heap(0, 0, Seq())), + heap.getOrElse(Heap(0, 0, Nil)), getGc(0, cs).toSeq ++ getGc(1, cs).toSeq ) } @@ -131,29 +159,17 @@ class Hotspot extends Jvm { private val log = Logger.getLogger(getClass.getName) private[this] val safepointBean = { - val runtimeBean = - try { - Class - .forName("sun.management.ManagementFactory") - .getMethod("getHotspotRuntimeMBean") - .invoke(null) - // jdk 6 has HotspotRuntimeMBean in the ManagementFactory class - } catch { - case _: Throwable => - Class - .forName("sun.management.ManagementFactoryHelper") - .getMethod("getHotspotRuntimeMBean") - .invoke(null) - // jdks 7 and 8 have HotspotRuntimeMBean in the ManagementFactoryHelper class - } + val runtimeBean = Class + .forName("sun.management.ManagementFactoryHelper") + .getMethod("getHotspotRuntimeMBean") + .invoke(null) def asSafepointBean(x: AnyRef) = { x.asInstanceOf[{ - def getSafepointSyncTime: Long; - def getTotalSafepointTime: Long; - def getSafepointCount: Long - } - ] + def getSafepointSyncTime: Long + def getTotalSafepointTime: Long + def getSafepointCount: Long + }] } try { asSafepointBean(runtimeBean) @@ -199,319 +215,67 @@ class Hotspot extends Jvm { def tenuringThreshold: Long = counter("sun.gc.policy.tenuringThreshold").map(long).getOrElse(0L) def snapCounters: Map[String, String] = - counters("") mapValues (_.getValue().toString) + counters("").map { case (k, v) => (k -> v.getValue().toString) } def forceGc(): Unit = System.gc() } -/* - -A sample of counters available on hotspot: - -scala> res0 foreach { case (k, v) => printf("%s = %s\n", k, v) } -java.ci.totalTime = 4453822466 -java.cls.loadedClasses = 3459 -java.cls.sharedLoadedClasses = 0 -java.cls.sharedUnloadedClasses = 0 -java.cls.unloadedClasses = 0 -java.property.java.class.path = . -java.property.java.endorsed.dirs = /System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Home/lib/endorsed -java.property.java.ext.dirs = /Library/Java/Extensions:/System/Library/Java/Extensions:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Home/lib/ext -java.property.java.home = /System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Home -java.property.java.library.path = /Applications/YourKit_Java_Profiler_9.5.6.app/bin/mac:.:/Library/Java/Extensions:/System/Library/Java/Extensions:/usr/lib/java -java.property.java.version = 1.6.0_29 -java.property.java.vm.info = mixed mode -java.property.java.vm.name = Java HotSpot(TM) 64-Bit Server VM -java.property.java.vm.specification.name = Java Virtual Machine Specification -java.property.java.vm.specification.vendor = Sun Microsystems Inc. -java.property.java.vm.specification.version = 1.0 -java.property.java.vm.vendor = Apple Inc. -java.property.java.vm.version = 20.4-b02-402 -java.rt.vmArgs = -Xserver -Djava.net.preferIPv4Stack=true -Dsbt.log.noformat=true -Xmx2G -XX:MaxPermSize=256m -XX:+UseConcMarkSweepGC -XX:+UseParNewGC -XX:MaxTenuringThreshold=3 -Xbootclasspath/a:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/jline.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-compiler.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-dbc.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-library.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-swing.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scalap.jar -Dscala.usejavacp=true -Dscala.home=/Users/marius/t/src/.scala/scala-2.8.1.final -Denv.emacs= -java.rt.vmFlags = -java.threads.daemon = 3 -java.threads.live = 4 -java.threads.livePeak = 5 -java.threads.started = 9 -sun.ci.compilerThread.0.compiles = 389 -sun.ci.compilerThread.0.method = -sun.ci.compilerThread.0.time = 52805 -sun.ci.compilerThread.0.type = 1 -sun.ci.compilerThread.1.compiles = 442 -sun.ci.compilerThread.1.method = java/lang/String compareTo -sun.ci.compilerThread.1.time = 49709 -sun.ci.compilerThread.1.type = 1 -sun.ci.lastFailedMethod = -sun.ci.lastFailedType = 0 -sun.ci.lastInvalidatedMethod = -sun.ci.lastInvalidatedType = 0 -sun.ci.lastMethod = java/nio/DirectByteBuffer ix -sun.ci.lastSize = 10 -sun.ci.lastType = 1 -sun.ci.nmethodCodeSize = 793568 -sun.ci.nmethodSize = 2272496 -sun.ci.osrBytes = 3470 -sun.ci.osrCompiles = 11 -sun.ci.osrTime = 75314565 -sun.ci.standardBytes = 262871 -sun.ci.standardCompiles = 819 -sun.ci.standardTime = 4394092506 -sun.ci.threads = 2 -sun.ci.totalBailouts = 0 -sun.ci.totalCompiles = 830 -sun.ci.totalInvalidates = 0 -sun.cls.appClassBytes = 174471 -sun.cls.appClassLoadCount = 165 -sun.cls.appClassLoadTime = 16577583 -sun.cls.appClassLoadTime.self = 8760586 -sun.cls.classInitTime = 153888518 -sun.cls.classInitTime.self = 75974400 -sun.cls.classLinkedTime = 76083753 -sun.cls.classLinkedTime.self = 61672107 -sun.cls.classVerifyTime = 14198241 -sun.cls.classVerifyTime.self = 8650545 -sun.cls.defineAppClassTime = 4361539 -sun.cls.defineAppClassTime.self = 579811 -sun.cls.defineAppClasses = 35 -sun.cls.initializedClasses = 2994 -sun.cls.isUnsyncloadClassSet = 0 -sun.cls.jniDefineClassNoLockCalls = 0 -sun.cls.jvmDefineClassNoLockCalls = 3 -sun.cls.jvmFindLoadedClassNoLockCalls = 0 -sun.cls.linkedClasses = 3353 -sun.cls.loadInstanceClassFailRate = 0 -sun.cls.loadedBytes = 9556504 -sun.cls.lookupSysClassTime = 165493723 -sun.cls.methodBytes = 7128160 -sun.cls.nonSystemLoaderLockContentionRate = 0 -sun.cls.parseClassTime = 260534425 -sun.cls.parseClassTime.self = 225060288 -sun.cls.sharedClassLoadTime = 259627 -sun.cls.sharedLoadedBytes = 0 -sun.cls.sharedUnloadedBytes = 0 -sun.cls.sysClassBytes = 16107656 -sun.cls.sysClassLoadTime = 518807017 -sun.cls.systemLoaderLockContentionRate = 0 -sun.cls.time = 547792239 -sun.cls.unloadedBytes = 0 -sun.cls.unsafeDefineClassCalls = 3 -sun.cls.verifiedClasses = 3353 -sun.gc.cause = No GC -sun.gc.collector.0.invocations = 13 -sun.gc.collector.0.lastEntryTime = 3235357990 -sun.gc.collector.0.lastExitTime = 3242811518 -sun.gc.collector.0.name = PCopy -sun.gc.collector.0.time = 74741924 -sun.gc.collector.1.invocations = 0 -sun.gc.collector.1.lastEntryTime = 0 -sun.gc.collector.1.lastExitTime = 0 -sun.gc.collector.1.name = CMS -sun.gc.collector.1.time = 0 -sun.gc.generation.0.agetable.bytes.00 = 0 -sun.gc.generation.0.agetable.bytes.01 = 2154984 -sun.gc.generation.0.agetable.bytes.02 = 0 -sun.gc.generation.0.agetable.bytes.03 = 0 -sun.gc.generation.0.agetable.bytes.04 = 0 -sun.gc.generation.0.agetable.bytes.05 = 0 -sun.gc.generation.0.agetable.bytes.06 = 0 -sun.gc.generation.0.agetable.bytes.07 = 0 -sun.gc.generation.0.agetable.bytes.08 = 0 -sun.gc.generation.0.agetable.bytes.09 = 0 -sun.gc.generation.0.agetable.bytes.10 = 0 -sun.gc.generation.0.agetable.bytes.11 = 0 -sun.gc.generation.0.agetable.bytes.12 = 0 -sun.gc.generation.0.agetable.bytes.13 = 0 -sun.gc.generation.0.agetable.bytes.14 = 0 -sun.gc.generation.0.agetable.bytes.15 = 0 -sun.gc.generation.0.agetable.size = 16 -sun.gc.generation.0.capacity = 21757952 -sun.gc.generation.0.maxCapacity = 174456832 -sun.gc.generation.0.minCapacity = 21757952 -sun.gc.generation.0.name = new -sun.gc.generation.0.space.0.capacity = 17432576 -sun.gc.generation.0.space.0.initCapacity = 0 -sun.gc.generation.0.space.0.maxCapacity = 139591680 -sun.gc.generation.0.space.0.name = eden -sun.gc.generation.0.space.0.used = 6954760 -sun.gc.generation.0.space.1.capacity = 2162688 -sun.gc.generation.0.space.1.initCapacity = 0 -sun.gc.generation.0.space.1.maxCapacity = 17432576 -sun.gc.generation.0.space.1.name = s0 -sun.gc.generation.0.space.1.used = 0 -sun.gc.generation.0.space.2.capacity = 2162688 -sun.gc.generation.0.space.2.initCapacity = 0 -sun.gc.generation.0.space.2.maxCapacity = 17432576 -sun.gc.generation.0.space.2.name = s1 -sun.gc.generation.0.space.2.used = 2162688 -sun.gc.generation.0.spaces = 3 -sun.gc.generation.0.threads = 8 -sun.gc.generation.1.capacity = 65404928 -sun.gc.generation.1.maxCapacity = 1973026816 -sun.gc.generation.1.minCapacity = 65404928 -sun.gc.generation.1.name = old -sun.gc.generation.1.space.0.capacity = 65404928 -sun.gc.generation.1.space.0.initCapacity = 65404928 -sun.gc.generation.1.space.0.maxCapacity = 1973026816 -sun.gc.generation.1.space.0.name = old -sun.gc.generation.1.space.0.used = 25122784 -sun.gc.generation.1.spaces = 1 -sun.gc.generation.2.capacity = 33161216 -sun.gc.generation.2.maxCapacity = 268435456 -sun.gc.generation.2.minCapacity = 21757952 -sun.gc.generation.2.name = perm -sun.gc.generation.2.space.0.capacity = 33161216 -sun.gc.generation.2.space.0.initCapacity = 21757952 -sun.gc.generation.2.space.0.maxCapacity = 268435456 -sun.gc.generation.2.space.0.name = perm -sun.gc.generation.2.space.0.used = 32439520 -sun.gc.generation.2.spaces = 1 -sun.gc.lastCause = unknown GCCause -sun.gc.policy.collectors = 2 -sun.gc.policy.desiredSurvivorSize = 1081344 -sun.gc.policy.generations = 3 -sun.gc.policy.maxTenuringThreshold = 3 -sun.gc.policy.name = ParNew:CMS -sun.gc.policy.tenuringThreshold = 1 -sun.gc.tlab.alloc = 756560 -sun.gc.tlab.allocThreads = 1 -sun.gc.tlab.fastWaste = 0 -sun.gc.tlab.fills = 28 -sun.gc.tlab.gcWaste = 0 -sun.gc.tlab.maxFastWaste = 0 -sun.gc.tlab.maxFills = 28 -sun.gc.tlab.maxGcWaste = 0 -sun.gc.tlab.maxSlowAlloc = 4 -sun.gc.tlab.maxSlowWaste = 116 -sun.gc.tlab.slowAlloc = 4 -sun.gc.tlab.slowWaste = 116 -sun.os.hrt.frequency = 1000000000 -sun.os.hrt.ticks = 6949680691 -sun.property.sun.boot.class.path = /System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/jsfd.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/classes.jar:/System/Library/Frameworks/JavaVM.framework/Frameworks/JavaRuntimeSupport.framework/Resources/Java/JavaRuntimeSupport.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/ui.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/laf.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/sunrsasign.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/jsse.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/jce.jar:/System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Classes/charsets.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/jline.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-compiler.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-dbc.jar:/Users/marius/t/src/.scala/scala-2.8.1.final/lib/scala-library.jar:/Users/marius/t/src/ -sun.property.sun.boot.library.path = /System/Library/Java/JavaVirtualMachines/1.6.0.jdk/Contents/Libraries -sun.rt._sync_ContendedLockAttempts = 53 -sun.rt._sync_Deflations = 11 -sun.rt._sync_EmptyNotifications = 0 -sun.rt._sync_FailedSpins = 0 -sun.rt._sync_FutileWakeups = 4 -sun.rt._sync_Inflations = 16 -sun.rt._sync_MonExtant = 256 -sun.rt._sync_MonInCirculation = 0 -sun.rt._sync_MonScavenged = 0 -sun.rt._sync_Notifications = 15 -sun.rt._sync_Parks = 43 -sun.rt._sync_PrivateA = 0 -sun.rt._sync_PrivateB = 0 -sun.rt._sync_SlowEnter = 0 -sun.rt._sync_SlowExit = 0 -sun.rt._sync_SlowNotify = 0 -sun.rt._sync_SlowNotifyAll = 0 -sun.rt._sync_SuccessfulSpins = 0 -sun.rt.applicationTime = 6804389569 -sun.rt.createVmBeginTime = 1320946078890 -sun.rt.createVmEndTime = 1320946078960 -sun.rt.internalVersion = Java HotSpot(TM) 64-Bit Server VM (20.4-b02-402) for macosx-amd64 JRE (1.6.0), built on Nov 1 2011 10:50:41 by "root" with gcc 4.2.1 (Based on Apple Inc. build 5658) (LLVM build 2335.15.00) -sun.rt.interruptedBeforeIO = 0 -sun.rt.interruptedDuringIO = 0 -sun.rt.javaCommand = scala.tools.nsc.MainGenericRunner -classpath /Users/marius/t/src/util/util-core/target/classes:/Users/marius/t/src/util/util-jvm/target/classes:/Users/marius/t/src/project/boot/scala-2.8.1/lib/scala-library.jar -sun.rt.jvmCapabilities = 1000000000000000000000000000000000000000000000000000000000000000 -sun.rt.jvmVersion = 335806466 -sun.rt.safepointSyncTime = 3107558 -sun.rt.safepointTime = 82605199 -sun.rt.safepoints = 21 -sun.rt.threadInterruptSignaled = 0 -sun.rt.vmInitDoneTime = 1320946078941 -sun.threads.vmOperationTime = 76932327 - - -## after a System.gc() - -sun.gc.cause = No GC -sun.gc.collector.0.invocations = 13 -sun.gc.collector.0.lastEntryTime = 3235357990 -sun.gc.collector.0.lastExitTime = 3242811518 -sun.gc.collector.0.name = PCopy -sun.gc.collector.0.time = 74741924 -sun.gc.collector.1.invocations = 0 -sun.gc.collector.1.lastEntryTime = 0 -sun.gc.collector.1.lastExitTime = 0 -sun.gc.collector.1.name = CMS -sun.gc.collector.1.time = 0 -sun.gc.generation.0.agetable.bytes.00 = 0 -sun.gc.generation.0.agetable.bytes.01 = 2154984 -sun.gc.generation.0.agetable.bytes.02 = 0 -sun.gc.generation.0.agetable.bytes.03 = 0 -sun.gc.generation.0.agetable.bytes.04 = 0 -sun.gc.generation.0.agetable.bytes.05 = 0 -sun.gc.generation.0.agetable.bytes.06 = 0 -sun.gc.generation.0.agetable.bytes.07 = 0 -sun.gc.generation.0.agetable.bytes.08 = 0 -sun.gc.generation.0.agetable.bytes.09 = 0 -sun.gc.generation.0.agetable.bytes.10 = 0 -sun.gc.generation.0.agetable.bytes.11 = 0 -sun.gc.generation.0.agetable.bytes.12 = 0 -sun.gc.generation.0.agetable.bytes.13 = 0 -sun.gc.generation.0.agetable.bytes.14 = 0 -sun.gc.generation.0.agetable.bytes.15 = 0 -sun.gc.generation.0.agetable.size = 16 -sun.gc.generation.0.capacity = 21757952 -sun.gc.generation.0.maxCapacity = 174456832 -sun.gc.generation.0.minCapacity = 21757952 -sun.gc.generation.0.name = new -sun.gc.generation.0.space.0.capacity = 17432576 -sun.gc.generation.0.space.0.initCapacity = 0 -sun.gc.generation.0.space.0.maxCapacity = 139591680 -sun.gc.generation.0.space.0.name = eden -sun.gc.generation.0.space.0.used = 6954760 -sun.gc.generation.0.space.1.capacity = 2162688 -sun.gc.generation.0.space.1.initCapacity = 0 -sun.gc.generation.0.space.1.maxCapacity = 17432576 -sun.gc.generation.0.space.1.name = s0 -sun.gc.generation.0.space.1.used = 0 -sun.gc.generation.0.space.2.capacity = 2162688 -sun.gc.generation.0.space.2.initCapacity = 0 -sun.gc.generation.0.space.2.maxCapacity = 17432576 -sun.gc.generation.0.space.2.name = s1 -sun.gc.generation.0.space.2.used = 2162688 -sun.gc.generation.0.spaces = 3 -sun.gc.generation.0.threads = 8 -sun.gc.generation.1.capacity = 65404928 -sun.gc.generation.1.maxCapacity = 1973026816 -sun.gc.generation.1.minCapacity = 65404928 -sun.gc.generation.1.name = old -sun.gc.generation.1.space.0.capacity = 65404928 -sun.gc.generation.1.space.0.initCapacity = 65404928 -sun.gc.generation.1.space.0.maxCapacity = 1973026816 -sun.gc.generation.1.space.0.name = old -sun.gc.generation.1.space.0.used = 25122784 -sun.gc.generation.1.spaces = 1 -sun.gc.generation.2.capacity = 33161216 -sun.gc.generation.2.maxCapacity = 268435456 -sun.gc.generation.2.minCapacity = 21757952 -sun.gc.generation.2.name = perm -sun.gc.generation.2.space.0.capacity = 33161216 -sun.gc.generation.2.space.0.initCapacity = 21757952 -sun.gc.generation.2.space.0.maxCapacity = 268435456 -sun.gc.generation.2.space.0.name = perm -sun.gc.generation.2.space.0.used = 32439520 -sun.gc.generation.2.spaces = 1 -sun.gc.lastCause = unknown GCCause -sun.gc.policy.collectors = 2 -sun.gc.policy.desiredSurvivorSize = 1081344 -sun.gc.policy.generations = 3 -sun.gc.policy.maxTenuringThreshold = 3 -sun.gc.policy.name = ParNew:CMS -sun.gc.policy.tenuringThreshold = 1 -sun.gc.tlab.alloc = 756560 -sun.gc.tlab.allocThreads = 1 -sun.gc.tlab.fastWaste = 0 -sun.gc.tlab.fills = 28 -sun.gc.tlab.gcWaste = 0 -sun.gc.tlab.maxFastWaste = 0 -sun.gc.tlab.maxFills = 28 -sun.gc.tlab.maxGcWaste = 0 -sun.gc.tlab.maxSlowAlloc = 4 -sun.gc.tlab.maxSlowWaste = 116 -sun.gc.tlab.slowAlloc = 4 -sun.gc.tlab.slowWaste = 116 - - */ +private object Hotspot { + + // this inlines the value of GarbageCollectionNotificationInfo.GARBAGE_COLLECTION_NOTIFICATION + // to avoid reflection + def isGcNotification(notification: Notification): Boolean = + notification.getType == "com.sun.management.gc.notification" + + private[this] val jvmStart = + Time.fromMilliseconds(ManagementFactory.getRuntimeMXBean.getStartTime) + + private[this] val gcNotifInfoFromMethod = + Class + .forName("com.sun.management.GarbageCollectionNotificationInfo") + .getMethod("from", classOf[CompositeData]) + + private[this] val gcNotifInfoGetGcInfoMethod = + Class + .forName("com.sun.management.GarbageCollectionNotificationInfo") + .getMethod("getGcInfo") + + private[this] val gcNotifInfoGetGcNameMethod = + Class + .forName("com.sun.management.GarbageCollectionNotificationInfo") + .getMethod("getGcName") + + private[this] val gcInfoGetIdMethod = + Class + .forName("com.sun.management.GcInfo") + .getMethod("getId") + + private[this] val gcInfoGetStartTimeMethod = + Class + .forName("com.sun.management.GcInfo") + .getMethod("getStartTime") + + private[this] val gcInfoGetDurationMethod = + Class + .forName("com.sun.management.GcInfo") + .getMethod("getDuration") + + /** + * This uses reflection because we the `com.sun` classes may not be available + * at compilation time. + */ + def gcFromNotificationInfo(compositeData: CompositeData): Gc = { + val gcNotifInfo /* com.sun.management.GarbageCollectionNotificationInfo */ = + gcNotifInfoFromMethod.invoke(null /* static method */, compositeData) + val gcInfo /* com.sun.management.GcInfo */ = + gcNotifInfoGetGcInfoMethod.invoke(gcNotifInfo) + Gc( + count = gcInfoGetIdMethod.invoke(gcInfo).asInstanceOf[Long], + name = gcNotifInfoGetGcNameMethod.invoke(gcNotifInfo).asInstanceOf[String], + timestamp = jvmStart + gcInfoGetStartTimeMethod + .invoke(gcInfo).asInstanceOf[Long].milliseconds, + duration = gcInfoGetDurationMethod.invoke(gcInfo).asInstanceOf[Long].milliseconds + ) + } + +} diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Jvm.scala b/util-jvm/src/main/scala/com/twitter/jvm/Jvm.scala index 5a8ba65a9a..5d9a157c76 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/Jvm.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/Jvm.scala @@ -1,97 +1,14 @@ package com.twitter.jvm import com.twitter.concurrent.NamedPoolThreadFactory -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util._ import java.lang.management.ManagementFactory import java.util.concurrent.{ConcurrentHashMap, Executors, ScheduledExecutorService, TimeUnit} import java.util.logging.{Level, Logger} -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ import scala.util.control.NonFatal -/** - * Information about the Heap - * - * @param allocated Estimated number of bytes - * that have been allocated so far (into eden) - * - * @param tenuringThreshold How many times an Object - * needs to be copied before being tenured. - * - * @param ageHisto Histogram of the number of bytes that - * have been copied as many times. Note: 0-indexed. - */ -case class Heap( - allocated: Long, - tenuringThreshold: Long, - ageHisto: Seq[Long] -) - -/** - * Information about the JVM's safepoint - * - * @param syncTimeMillis Cumulative time, in milliseconds, spent - * getting all threads to safepoint states - * - * @param totalTimeMillis Cumulative time, in milliseconds, that the - * application has been stopped for safepoint operations - * - * @param count The number of safepoints taken place since - * the JVM started - */ -case class Safepoint(syncTimeMillis: Long, totalTimeMillis: Long, count: Long) - -case class PoolState(numCollections: Long, capacity: StorageUnit, used: StorageUnit) { - def -(other: PoolState) = PoolState( - numCollections = this.numCollections - other.numCollections, - capacity = other.capacity, - used = this.used + other.capacity - other.used + - other.capacity * (this.numCollections - other.numCollections - 1) - ) - - override def toString = - "PoolState(n=%d,remaining=%s[%s of %s])".format(numCollections, capacity - used, used, capacity) -} - -/** - * A handle to a garbage collected memory pool. - */ -trait Pool { - - /** Get the current state of this memory pool. */ - def state(): PoolState - - /** - * Sample the allocation rate of this pool. Note that this is merely - * an estimation based on sampling the state of the pool initially - * and then again when the period elapses. - * - * @return Future of the samples rate (in bps). - */ - def estimateAllocRate(period: Duration, timer: Timer): Future[Long] = { - val elapsed = Stopwatch.start() - val begin = state() - timer.doLater(period) { - val end = state() - val interval = elapsed() - ((end - begin).used.inBytes * 1000) / interval.inMilliseconds - } - } -} - -case class Gc( - count: Long, - name: String, - timestamp: Time, - duration: Duration -) - -case class Snapshot( - timestamp: Time, - heap: Heap, - lastGcs: Seq[Gc] -) - /** * Access JVM internal performance counters. We maintain a strict * interface so that we are decoupled from the actual underlying JVM. @@ -99,6 +16,9 @@ case class Snapshot( trait Jvm { import Jvm.log + // This is intended to only be overridden by tests + protected def logger: Logger = log + trait Opts { def compileThresh: Option[Int] } @@ -118,7 +38,7 @@ trait Jvm { */ def snap: Snapshot - def edenPool: Pool + def edenPool: com.twitter.jvm.Pool /** * Gets the current usage of the metaspace, if available. @@ -170,8 +90,8 @@ trait Jvm { } else if (lastCount != count) { missedCollections += count - 1 - lastCount if (missedCollections > 0 && Time.now - lastLog > LogPeriod) { - if (log.isLoggable(Level.FINE)) { - log.fine( + if (logger.isLoggable(Level.FINE)) { + logger.fine( "Missed %d collections for %s due to sampling" .format(missedCollections, name) ) @@ -187,7 +107,7 @@ trait Jvm { } executor.scheduleAtFixedRate( - new Runnable { def run() = sample() }, + new Runnable { def run(): Unit = sample() }, 0 /*initial delay*/, Period.inMilliseconds, TimeUnit.MILLISECONDS @@ -211,8 +131,7 @@ trait Jvm { buffer = (gc :: buffer).takeWhile(_.timestamp > floor) } - (since: Time) => - buffer.takeWhile(_.timestamp > since) + (since: Time) => buffer.takeWhile(_.timestamp > since) } def forceGc(): Unit @@ -226,9 +145,9 @@ trait Jvm { */ def mainClassName: String = { val mainClass = for { - (_, stack) <- Thread.getAllStackTraces().asScala.find { case (t, s) => t.getName == "main" } + (_, stack) <- Thread.getAllStackTraces.asScala.find { case (t, s) => t.getName == "main" } frame <- stack.reverse.find { elem => - !(elem.getClassName.startsWith("scala.tools.nsc.MainGenericRunner")) + !elem.getClassName.startsWith("scala.tools.nsc.MainGenericRunner") } } yield frame.getClassName @@ -271,7 +190,7 @@ object Jvm { /** * Return an instance of the [[Jvm]] for this runtime. * - * See [[Jvms.apply()]] for Java compatibility. + * See `Jvms.apply` for Java compatibility. */ def apply(): Jvm = _jvm @@ -299,7 +218,7 @@ object Jvms { Jvm.ProcessId /** - * Java compatibility for [[Jvm.apply()]]. + * Java compatibility for [[Jvm.apply]]. */ def get(): Jvm = Jvm() diff --git a/util-jvm/src/main/scala/com/twitter/jvm/JvmStats.scala b/util-jvm/src/main/scala/com/twitter/jvm/JvmStats.scala index c8bad9d714..b167903e82 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/JvmStats.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/JvmStats.scala @@ -1,10 +1,21 @@ package com.twitter.jvm -import com.twitter.conversions.string._ -import com.twitter.finagle.stats.{BroadcastStatsReceiver, StatsReceiver} -import java.lang.management.{ManagementFactory, BufferPoolMXBean} -import scala.collection.JavaConverters._ +import com.sun.management.GarbageCollectionNotificationInfo +import com.twitter.conversions.StringOps._ +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.Bytes +import com.twitter.finagle.stats.Milliseconds +import com.twitter.finagle.stats.StatsReceiver +import com.twitter.finagle.stats.exp.Expression +import com.twitter.finagle.stats.exp.ExpressionSchema +import java.lang.management.BufferPoolMXBean +import java.lang.management.ManagementFactory +import javax.management.Notification +import javax.management.NotificationEmitter +import javax.management.NotificationListener +import javax.management.openmbean.CompositeData import scala.collection.mutable +import scala.jdk.CollectionConverters._ object JvmStats { // set used for keeping track of jvm gauges (otherwise only weakly referenced) @@ -20,76 +31,115 @@ object JvmStats { def heap = mem.getHeapMemoryUsage() val heapStats = stats.scope("heap") - gauges.add(heapStats.addGauge("committed") { heap.getCommitted() }) - gauges.add(heapStats.addGauge("max") { heap.getMax() }) - gauges.add(heapStats.addGauge("used") { heap.getUsed() }) + val heapUsedGauge = heapStats.addGauge("used") { heap.getUsed().toFloat } + gauges.add(heapStats.addGauge("committed") { heap.getCommitted().toFloat }) + gauges.add(heapStats.addGauge("max") { heap.getMax().toFloat }) + gauges.add(heapUsedGauge) def nonHeap = mem.getNonHeapMemoryUsage() val nonHeapStats = stats.scope("nonheap") - gauges.add(nonHeapStats.addGauge("committed") { nonHeap.getCommitted() }) - gauges.add(nonHeapStats.addGauge("max") { nonHeap.getMax() }) - gauges.add(nonHeapStats.addGauge("used") { nonHeap.getUsed() }) + gauges.add(nonHeapStats.addGauge("committed") { nonHeap.getCommitted().toFloat }) + gauges.add(nonHeapStats.addGauge("max") { nonHeap.getMax().toFloat }) + gauges.add(nonHeapStats.addGauge("used") { nonHeap.getUsed().toFloat }) val threads = ManagementFactory.getThreadMXBean() val threadStats = stats.scope("thread") - gauges.add(threadStats.addGauge("daemon_count") { threads.getDaemonThreadCount().toLong }) - gauges.add(threadStats.addGauge("count") { threads.getThreadCount().toLong }) - gauges.add(threadStats.addGauge("peak_count") { threads.getPeakThreadCount().toLong }) + gauges.add(threadStats.addGauge("daemon_count") { threads.getDaemonThreadCount().toFloat }) + gauges.add(threadStats.addGauge("count") { threads.getThreadCount().toFloat }) + gauges.add(threadStats.addGauge("peak_count") { threads.getPeakThreadCount().toFloat }) val runtime = ManagementFactory.getRuntimeMXBean() - gauges.add(stats.addGauge("start_time") { runtime.getStartTime() }) - gauges.add(stats.addGauge("uptime") { runtime.getUptime() }) + val uptime = stats.addGauge("uptime") { runtime.getUptime().toFloat } + gauges.add(uptime) + gauges.add(stats.addGauge("start_time") { runtime.getStartTime().toFloat }) + gauges.add(stats.addGauge("spec_version") { runtime.getSpecVersion.toFloat }) val os = ManagementFactory.getOperatingSystemMXBean() - gauges.add(stats.addGauge("num_cpus") { os.getAvailableProcessors().toLong }) + gauges.add(stats.addGauge("num_cpus") { os.getAvailableProcessors().toFloat }) os match { case unix: com.sun.management.UnixOperatingSystemMXBean => - gauges.add(stats.addGauge("fd_count") { unix.getOpenFileDescriptorCount }) - gauges.add(stats.addGauge("fd_limit") { unix.getMaxFileDescriptorCount }) + val fileDescriptorCountGauge = stats.addGauge("fd_count") { + unix.getOpenFileDescriptorCount.toFloat + } + gauges.add(fileDescriptorCountGauge) + gauges.add(stats.addGauge("fd_limit") { unix.getMaxFileDescriptorCount.toFloat }) + // register expression + stats.registerExpression( + ExpressionSchema("file_descriptors", Expression(fileDescriptorCountGauge.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withDescription( + "Total file descriptors used by the service. If it continuously increasing over time, then potentially files or connections aren't being closed")) case _ => } + val compilerStats = stats.scope("compiler"); + gauges.add(compilerStats.addGauge("graal") { + System.getProperty("jvmci.Compiler") match { + case "graal" => 1 + case _ => 0 + } + }) + ManagementFactory.getCompilationMXBean() match { case null => case compilation => val compilationStats = stats.scope("compilation") - gauges.add(compilationStats.addGauge("time_msec") { compilation.getTotalCompilationTime() }) + gauges.add( + compilationStats.addGauge( + compilationStats.metricBuilder(GaugeType).withCounterishGauge.withName("time_msec")) { + compilation.getTotalCompilationTime().toFloat + }) } val classes = ManagementFactory.getClassLoadingMXBean() val classLoadingStats = stats.scope("classes") - gauges.add(classLoadingStats.addGauge("total_loaded") { classes.getTotalLoadedClassCount() }) - gauges.add(classLoadingStats.addGauge("total_unloaded") { classes.getUnloadedClassCount() }) gauges.add( - classLoadingStats.addGauge("current_loaded") { classes.getLoadedClassCount().toLong } + classLoadingStats.addGauge( + classLoadingStats.metricBuilder(GaugeType).withCounterishGauge.withName("total_loaded")) { + classes.getTotalLoadedClassCount().toFloat + }) + gauges.add( + classLoadingStats.addGauge( + classLoadingStats.metricBuilder(GaugeType).withCounterishGauge.withName("total_unloaded")) { + classes.getUnloadedClassCount().toFloat + }) + gauges.add( + classLoadingStats.addGauge("current_loaded") { classes.getLoadedClassCount().toFloat } ) val memPool = ManagementFactory.getMemoryPoolMXBeans.asScala val memStats = stats.scope("mem") val currentMem = memStats.scope("current") - // TODO: Refactor postGCStats when we confirmed that no one is using this stats anymore - // val postGCStats = memStats.scope("postGC") - val postGCMem = memStats.scope("postGC") - val postGCStats = BroadcastStatsReceiver(Seq(stats.scope("postGC"), postGCMem)) + val postGCStats = memStats.scope("postGC") memPool.foreach { pool => - val name = pool.getName.regexSub("""[^\w]""".r) { m => - "_" - } + val name = pool.getName.regexSub("""[^\w]""".r) { m => "_" } + val currentPoolScope = currentMem.hierarchicalScope(name).label("pool", name) + val postGcPoolScope = postGCStats.hierarchicalScope(name).label("pool", name) if (pool.getCollectionUsage != null) { def usage = pool.getCollectionUsage // this is a snapshot, we can't reuse the value - gauges.add(postGCStats.addGauge(name, "used") { usage.getUsed }) + gauges.add(postGcPoolScope.addGauge("used") { usage.getUsed.toFloat }) } if (pool.getUsage != null) { def usage = pool.getUsage // this is a snapshot, we can't reuse the value - gauges.add(currentMem.addGauge(name, "used") { usage.getUsed }) - gauges.add(currentMem.addGauge(name, "max") { usage.getMax }) + val usageGauge = currentPoolScope.addGauge("used") { usage.getUsed.toFloat } + gauges.add(usageGauge) + gauges.add(currentPoolScope.addGauge("max") { usage.getMax.toFloat }) + + // register memory usage expression + currentMem.registerExpression( + ExpressionSchema("memory_pool", Expression(usageGauge.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withLabel("kind", name) + .withUnit(Bytes) + .withDescription( + s"The current estimate of the amount of space within the $name memory pool holding allocated objects in bytes")) } } gauges.add(postGCStats.addGauge("used") { - memPool.flatMap(p => Option(p.getCollectionUsage)).map(_.getUsed).sum + memPool.flatMap(p => Option(p.getCollectionUsage)).map(_.getUsed).sum.toFloat }) gauges.add(currentMem.addGauge("used") { - memPool.flatMap(p => Option(p.getUsage)).map(_.getUsed).sum + memPool.flatMap(p => Option(p.getUsage)).map(_.getUsed).sum.toFloat }) // the Hotspot JVM exposes the full size that the metaspace can grow to @@ -97,14 +147,25 @@ object JvmStats { val jvm = Jvm() jvm.metaspaceUsage.foreach { usage => gauges.add(memStats.scope("metaspace").addGauge("max_capacity") { - usage.maxCapacity.inBytes + usage.maxCapacity.inBytes.toFloat }) } val spStats = stats.scope("safepoint") - gauges.add(spStats.addGauge("sync_time_millis") { jvm.safepoint.syncTimeMillis }) - gauges.add(spStats.addGauge("total_time_millis") { jvm.safepoint.totalTimeMillis }) - gauges.add(spStats.addGauge("count") { jvm.safepoint.count }) + gauges.add( + spStats.addGauge( + spStats.metricBuilder(GaugeType).withCounterishGauge.withName("sync_time_millis")) { + jvm.safepoint.syncTimeMillis.toFloat + }) + gauges.add( + spStats.addGauge( + spStats.metricBuilder(GaugeType).withCounterishGauge.withName("total_time_millis")) { + jvm.safepoint.totalTimeMillis.toFloat + }) + gauges.add( + spStats.addGauge(spStats.metricBuilder(GaugeType).withCounterishGauge.withName("count")) { + jvm.safepoint.count.toFloat + }) ManagementFactory.getPlatformMXBeans(classOf[BufferPoolMXBean]) match { case null => @@ -112,37 +173,109 @@ object JvmStats { val bufferPoolStats = memStats.scope("buffer") jBufferPool.asScala.foreach { bp => val name = bp.getName - gauges.add(bufferPoolStats.addGauge(name, "count") { bp.getCount }) - gauges.add(bufferPoolStats.addGauge(name, "used") { bp.getMemoryUsed }) - gauges.add(bufferPoolStats.addGauge(name, "max") { bp.getTotalCapacity }) + val bufferPoolScope = bufferPoolStats.hierarchicalScope(name).label("pool", name) + gauges.add(bufferPoolScope.addGauge("count") { bp.getCount.toFloat }) + gauges.add(bufferPoolScope.addGauge("used") { bp.getMemoryUsed.toFloat }) + gauges.add(bufferPoolScope.addGauge("max") { bp.getTotalCapacity.toFloat }) } } val gcPool = ManagementFactory.getGarbageCollectorMXBeans.asScala val gcStats = stats.scope("gc") gcPool.foreach { gc => - val name = gc.getName.regexSub("""[^\w]""".r) { m => - "_" + val name = gc.getName.regexSub("""[^\w]""".r) { m => "_" } + val poolScope = gcStats.hierarchicalScope(name).label("jmm", name) + val poolCycles = + poolScope.addGauge( + poolScope.metricBuilder(GaugeType).withCounterishGauge.withName("cycles")) { + gc.getCollectionCount.toFloat + } + val poolMsec = poolScope.addGauge( + poolScope.metricBuilder(GaugeType).withCounterishGauge.withName("msec")) { + gc.getCollectionTime.toFloat } - gauges.add(gcStats.addGauge(name, "cycles") { gc.getCollectionCount }) - gauges.add(gcStats.addGauge(name, "msec") { gc.getCollectionTime }) + + val gcPauseStat = poolScope.stat("pause_msec") + gc.asInstanceOf[NotificationEmitter].addNotificationListener( + new NotificationListener { + override def handleNotification( + notification: Notification, + handback: Any + ): Unit = { + notification.getType + if (notification.getType == GarbageCollectionNotificationInfo.GARBAGE_COLLECTION_NOTIFICATION) { + gcPauseStat.add( + GarbageCollectionNotificationInfo + .from(notification.getUserData + .asInstanceOf[CompositeData]).getGcInfo.getDuration) + } + } + }, + null, + null + ) + + gcStats.registerExpression( + ExpressionSchema(s"gc_cycles", Expression(poolCycles.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withLabel("gc_pool", name) + .withDescription( + s"The total number of collections that have occurred for the $name gc pool")) + gcStats.registerExpression(ExpressionSchema(s"gc_latency", Expression(poolMsec.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withLabel("gc_pool", name) + .withUnit(Milliseconds) + .withDescription( + s"The total elapsed time spent doing collections for the $name gc pool in milliseconds")) + + gauges.add(poolCycles) + gauges.add(poolMsec) } // note, these could be -1 if the collector doesn't have support for it. - gauges.add(gcStats.addGauge("cycles") { gcPool.map(_.getCollectionCount).filter(_ > 0).sum }) - gauges.add(gcStats.addGauge("msec") { gcPool.map(_.getCollectionTime).filter(_ > 0).sum }) + val cycles = + gcStats.addGauge(gcStats.metricBuilder(GaugeType).withCounterishGauge.withName("cycles")) { + gcPool.map(_.getCollectionCount).filter(_ > 0).sum.toFloat + } + val msec = + gcStats.addGauge(gcStats.metricBuilder(GaugeType).withCounterishGauge.withName("msec")) { + gcPool.map(_.getCollectionTime).filter(_ > 0).sum.toFloat + } + gauges.add(cycles) + gauges.add(msec) allocations = new Allocations(gcStats) allocations.start() if (allocations.trackingEden) { val allocationStats = memStats.scope("allocations") - val eden = allocationStats.scope("eden") - gauges.add(eden.addGauge("bytes") { allocations.eden }) + val eden = allocationStats.hierarchicalScope("eden").label("space", "eden") + gauges.add(eden.addGauge("bytes") { allocations.eden.toFloat }) } // return ms from ns while retaining precision gauges.add(stats.addGauge("application_time_millis") { jvm.applicationTime.toFloat / 1000000 }) - gauges.add(stats.addGauge("tenuring_threshold") { jvm.tenuringThreshold }) - } + gauges.add(stats.addGauge("tenuring_threshold") { jvm.tenuringThreshold.toFloat }) + // register metric expressions + stats.registerExpression( + ExpressionSchema("jvm_uptime", Expression(uptime.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withUnit(Milliseconds) + .withDescription("The uptime of the JVM in milliseconds")) + gcStats.registerExpression( + ExpressionSchema("gc_cycles", Expression(cycles.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withDescription("The total number of collections that have occurred")) + gcStats.registerExpression( + ExpressionSchema("gc_latency", Expression(msec.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withUnit(Milliseconds) + .withDescription("The total elapsed time spent doing collections in milliseconds")) + heapStats.registerExpression( + ExpressionSchema("memory_pool", Expression(heapUsedGauge.metadata)) + .withLabel(ExpressionSchema.Role, "jvm") + .withLabel("kind", "Heap") + .withUnit(Bytes) + .withDescription("Heap in use in bytes")) + } } diff --git a/util-jvm/src/main/scala/com/twitter/jvm/NilJvm.scala b/util-jvm/src/main/scala/com/twitter/jvm/NilJvm.scala index f98e07f479..a781fedd9e 100644 --- a/util-jvm/src/main/scala/com/twitter/jvm/NilJvm.scala +++ b/util-jvm/src/main/scala/com/twitter/jvm/NilJvm.scala @@ -1,11 +1,11 @@ package com.twitter.jvm -import com.twitter.conversions.storage._ +import com.twitter.conversions.StorageUnitOps._ import com.twitter.util.Time object NilJvm extends Jvm { val opts: Opts = new Opts { - def compileThresh = None + def compileThresh: Option[Int] = None } def forceGc(): Unit = System.gc() val snapCounters: Map[String, String] = Map() diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Opts.scala b/util-jvm/src/main/scala/com/twitter/jvm/Opt.scala similarity index 100% rename from util-jvm/src/main/scala/com/twitter/jvm/Opts.scala rename to util-jvm/src/main/scala/com/twitter/jvm/Opt.scala diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Pool.scala b/util-jvm/src/main/scala/com/twitter/jvm/Pool.scala new file mode 100644 index 0000000000..dff4649def --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/Pool.scala @@ -0,0 +1,29 @@ +package com.twitter.jvm + +import com.twitter.util.{Duration, Future, Stopwatch, Timer} + +/** + * A handle to a garbage collected memory pool. + */ +trait Pool { + + /** Get the current state of this memory pool. */ + def state(): PoolState + + /** + * Sample the allocation rate of this pool. Note that this is merely + * an estimation based on sampling the state of the pool initially + * and then again when the period elapses. + * + * @return Future of the samples rate (in bps). + */ + def estimateAllocRate(period: Duration, timer: Timer): Future[Long] = { + val elapsed = Stopwatch.start() + val begin = state() + timer.doLater(period) { + val end = state() + val interval = elapsed() + ((end - begin).used.inBytes * 1000) / interval.inMilliseconds + } + } +} diff --git a/util-jvm/src/main/scala/com/twitter/jvm/PoolState.scala b/util-jvm/src/main/scala/com/twitter/jvm/PoolState.scala new file mode 100644 index 0000000000..5d02289be8 --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/PoolState.scala @@ -0,0 +1,15 @@ +package com.twitter.jvm + +import com.twitter.util.StorageUnit + +case class PoolState(numCollections: Long, capacity: StorageUnit, used: StorageUnit) { + def -(other: PoolState): PoolState = PoolState( + numCollections = this.numCollections - other.numCollections, + capacity = other.capacity, + used = this.used + other.capacity - other.used + + other.capacity * (this.numCollections - other.numCollections - 1) + ) + + override def toString: String = + s"PoolState(n=$numCollections,remaining=${capacity - used}[$used of $capacity])" +} diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Safepoint.scala b/util-jvm/src/main/scala/com/twitter/jvm/Safepoint.scala new file mode 100644 index 0000000000..f08bed510e --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/Safepoint.scala @@ -0,0 +1,15 @@ +package com.twitter.jvm + +/** + * Information about the JVM's safepoint + * + * @param syncTimeMillis Cumulative time, in milliseconds, spent + * getting all threads to safepoint states + * + * @param totalTimeMillis Cumulative time, in milliseconds, that the + * application has been stopped for safepoint operations + * + * @param count The number of safepoints taken place since + * the JVM started + */ +case class Safepoint(syncTimeMillis: Long, totalTimeMillis: Long, count: Long) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Snapshot.scala b/util-jvm/src/main/scala/com/twitter/jvm/Snapshot.scala new file mode 100644 index 0000000000..642d6e090b --- /dev/null +++ b/util-jvm/src/main/scala/com/twitter/jvm/Snapshot.scala @@ -0,0 +1,5 @@ +package com.twitter.jvm + +import com.twitter.util.Time + +case class Snapshot(timestamp: Time, heap: Heap, lastGcs: Seq[Gc]) diff --git a/util-jvm/src/main/scala/com/twitter/jvm/Cpu.scala b/util-jvm/src/main/scala/com/twitter/jvm/numProcs.scala similarity index 100% rename from util-jvm/src/main/scala/com/twitter/jvm/Cpu.scala rename to util-jvm/src/main/scala/com/twitter/jvm/numProcs.scala diff --git a/util-jvm/src/test/java/BUILD b/util-jvm/src/test/java/BUILD index 0a56246fcd..32ca28064d 100644 --- a/util-jvm/src/test/java/BUILD +++ b/util-jvm/src/test/java/BUILD @@ -1,8 +1,15 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "util/util-jvm/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-jvm/src/main/scala/com/twitter/jvm", ], ) diff --git a/util-jvm/src/test/scala/BUILD b/util-jvm/src/test/scala/BUILD index ca099b7fb2..955928957e 100644 --- a/util-jvm/src/test/scala/BUILD +++ b/util-jvm/src/test/scala/BUILD @@ -1,22 +1,51 @@ +COMMON_DEPS = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:mockito", + "util/util-app/src/main/scala", + "util/util-core:util-core-util", + "util/util-jvm/src/main/scala/com/twitter/jvm", +] + junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-jvm/src/main/scala", - "util/util-logging/src/main/scala", - "util/util-test/src/main/scala", + sources = [ + "!**/EstimatorApp.scala", + "**/*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = COMMON_DEPS + [ + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-stats/src/main/scala/com/twitter/finagle/stats:stats", ], ) +scala_library( + name = "estimator_app", + # Overlap the sources (except EstimatorApp.scala) with the above target for the following reasons: + # 1) We cannot extract a common `scala_library` because Pants does not like `junit_tests` + # depending on a `scala_library` containing the test sources while the `junit_test` + # itself has empty sources. + # 2) Bazel strictly forbids `scala_library` to depend back onto a test target, whereas Pants + # does not strictly enforce it. + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = COMMON_DEPS + ["util/util-stats/src/main/scala/com/twitter/finagle/stats:stats"], +) + jvm_binary( name = "bin", main = "com.twitter.jvm.EstimatorApp", - dependencies = [ - ":scala", - ], + platform = "java8", + tags = ["bazel-compatible"], + dependencies = [":estimator_app"], ) diff --git a/util-jvm/src/test/scala/com/twitter/jvm/ContentionTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/ContentionTest.scala index f5cd2c8b7d..fe2305efd0 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/ContentionTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/ContentionTest.scala @@ -1,29 +1,24 @@ package com.twitter.jvm +import com.twitter.util.{Await, Promise} import java.util.concurrent.locks.ReentrantLock - -import org.junit.runner.RunWith -import org.scalatest.FunSuite import org.scalatest.concurrent.Eventually -import org.scalatest.junit.JUnitRunner import org.scalatest.time.{Millis, Seconds, Span} - -import com.twitter.util.{Await, Promise} +import org.scalatest.funsuite.AnyFunSuite class Philosopher { val ready = new Promise[Unit] private val lock = new ReentrantLock() def dine(neighbor: Philosopher): Unit = { lock.lockInterruptibly() - ready.setValue(Unit) + ready.setValue(()) Await.ready(neighbor.ready) neighbor.dine(this) lock.unlock() } } -@RunWith(classOf[JUnitRunner]) -class ContentionTest extends FunSuite with Eventually { +class ContentionTest extends AnyFunSuite with Eventually { implicit override val patienceConfig = PatienceConfig(timeout = scaled(Span(15, Seconds)), interval = scaled(Span(5, Millis))) diff --git a/util-jvm/src/test/scala/com/twitter/jvm/CpuProfileTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/CpuProfileTest.scala index ae255d9eac..c806175927 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/CpuProfileTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/CpuProfileTest.scala @@ -1,21 +1,19 @@ package com.twitter.jvm -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.Time import java.io.ByteArrayOutputStream -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.compat.immutable.LazyList -@RunWith(classOf[JUnitRunner]) -class CpuProfileTest extends FunSuite { +class CpuProfileTest extends AnyFunSuite { test("record") { // record() calls Time.now 3 times initially, and then 3 times on every loop iteration. - val times: Stream[Int] = (0 #:: Stream.from(0)).flatMap(x => List(x, x, x)) + val times: LazyList[Int] = (0 #:: LazyList.from(0)).flatMap(x => List(x, x, x)) val iter = times.iterator val start = Time.now - def nextTime: Time = start + iter.next().milliseconds * 10 + def nextTime: Time = start + iter.next().toLong.milliseconds * 10 val t = new Thread("CpuProfileTest") { override def run(): Unit = { diff --git a/util-jvm/src/test/scala/com/twitter/jvm/EstimatorApp.scala b/util-jvm/src/test/scala/com/twitter/jvm/EstimatorApp.scala new file mode 100644 index 0000000000..8ac1c40084 --- /dev/null +++ b/util-jvm/src/test/scala/com/twitter/jvm/EstimatorApp.scala @@ -0,0 +1,93 @@ +package com.twitter.jvm + +/** + * Take a GC log produced by: + * + * {{{ + * $ jstat -gc \$PID 250 ... + * }}} + * + * And report on GC prediction accuracy. Time is + * indexed by the jstat output, and the columns are, + * in order: current time, next gc, estimated next GC. + */ +object EstimatorApp extends App { + import com.twitter.conversions.StorageUnitOps._ + + val estimator = args match { + case Array("kalman", n, error) => + new KalmanGaussianError(n.toInt, error.toDouble) + case Array("windowed", n, windows) => + new WindowedMeans( + n.toInt, + windows + .split(",").iterator.map { w => + w.split(":") match { + case Array(w, i) => (w.toInt, i.toInt) + case _ => throw new IllegalArgumentException("bad weight, count pair " + w) + } + }.toSeq + ) + case Array("load", interval) => + new LoadAverage(interval.toDouble) + case _ => throw new IllegalArgumentException("bad args ") + } + + val lines = scala.io.Source.stdin.getLines().drop(1) + val states = lines.toArray map (_.split(" ") filter (_ != "") map (_.toDouble)) collect { + case Array(s0c, s1c, s0u, s1u, ec, eu, oc, ou, pc, pu, ygc, ygct, fgc, fgct, gct) => + PoolState(ygc.toLong, ec.toLong.bytes, eu.toLong.bytes) + } + + var elapsed = 1 + for (List(begin, end) <- states.toList.sliding(2)) { + val allocated = (end - begin).used + estimator.measure(allocated.inBytes.toDouble) + val r = end.capacity - end.used + val i = (r.inBytes / estimator.estimate.toLong) + elapsed + val j = states.indexWhere(_.numCollections > end.numCollections) + + if (j > 0) + println("%d %d %d".format(elapsed, j, i)) + + elapsed += 1 + } +} + +/* + +The following script is useful for plotting +results from EstimatorApp: + + % scala ... com.twitter.jvm.EstimatorApp [ARGS] > /tmp/out + +#!/usr/bin/env gnuplot + +set terminal png size 800,600 +set title "GC predictor" + +set macros + +set grid +set timestamp "Generated on %Y-%m-%d by `whoami`" font "Helvetica-Oblique, 8pt" +set noclip +set xrange [0:1100] + +set key reverse left Left # box + +set ylabel "Time of GC" textcolor lt 3 +set yrange [0:1100] +set mytics 5 +set y2tics + +set xlabel "Time" textcolor lt 4 +set mxtics 5 + +set boxwidth 0.5 +set style fill transparent pattern 4 bo + +plot "< awk '{print $1 \" \" $2}' /tmp/out" title "Actual", \ + "< awk '{print $1 \" \" $3}' /tmp/out" title "Predicted", \ + "< awk '{print $1 \" \" $1}' /tmp/out" title "time" with lines + + */ diff --git a/util-jvm/src/test/scala/com/twitter/jvm/EstimatorTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/EstimatorTest.scala index 96b849e6b4..95c8b20e02 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/EstimatorTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/EstimatorTest.scala @@ -1,14 +1,11 @@ package com.twitter.jvm -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class EstimatorTest extends FunSuite { +class EstimatorTest extends AnyFunSuite { test("LoadAverage") { // This makes LoadAverage.a = 1/2 for easy testing. - val interval = -1D / math.log(0.5) + val interval = -1d / math.log(0.5) val e = new LoadAverage(interval) assert(e.estimate.isNaN) @@ -22,94 +19,3 @@ class EstimatorTest extends FunSuite { assert(e.estimate == 0) } } - -/** - * Take a GC log produced by: - * - * {{{ - * $ jstat -gc \$PID 250 ... - * }}} - * - * And report on GC prediction accuracy. Time is - * indexed by the jstat output, and the columns are, - * in order: current time, next gc, estimated next GC. - */ -object EstimatorApp extends App { - import com.twitter.conversions.storage._ - - val estimator = args match { - case Array("kalman", n, error) => - new KalmanGaussianError(n.toInt, error.toDouble) - case Array("windowed", n, windows) => - new WindowedMeans( - n.toInt, - windows.split(",") map { w => - w.split(":") match { - case Array(w, i) => (w.toInt, i.toInt) - case _ => throw new IllegalArgumentException("bad weight, count pair " + w) - } - } - ) - case Array("load", interval) => - new LoadAverage(interval.toDouble) - case _ => throw new IllegalArgumentException("bad args ") - } - - val lines = scala.io.Source.stdin.getLines().drop(1) - val states = lines.toArray map (_.split(" ") filter (_ != "") map (_.toDouble)) collect { - case Array(s0c, s1c, s0u, s1u, ec, eu, oc, ou, pc, pu, ygc, ygct, fgc, fgct, gct) => - PoolState(ygc.toLong, ec.toLong.bytes, eu.toLong.bytes) - } - - var elapsed = 1 - for (List(begin, end) <- states.toList.sliding(2)) { - val allocated = (end - begin).used - estimator.measure(allocated.inBytes) - val r = end.capacity - end.used - val i = (r.inBytes / estimator.estimate.toLong) + elapsed - val j = states.indexWhere(_.numCollections > end.numCollections) - - if (j > 0) - println("%d %d %d".format(elapsed, j, i)) - - elapsed += 1 - } -} - -/* - -The following script is useful for plotting -results from EstimatorApp: - - % scala ... com.twitter.jvm.EstimatorApp [ARGS] > /tmp/out - -#!/usr/bin/env gnuplot - -set terminal png size 800,600 -set title "GC predictor" - -set macros - -set grid -set timestamp "Generated on %Y-%m-%d by `whoami`" font "Helvetica-Oblique, 8pt" -set noclip -set xrange [0:1100] - -set key reverse left Left # box - -set ylabel "Time of GC" textcolor lt 3 -set yrange [0:1100] -set mytics 5 -set y2tics - -set xlabel "Time" textcolor lt 4 -set mxtics 5 - -set boxwidth 0.5 -set style fill transparent pattern 4 bo - -plot "< awk '{print $1 \" \" $2}' /tmp/out" title "Actual", \ - "< awk '{print $1 \" \" $3}' /tmp/out" title "Predicted", \ - "< awk '{print $1 \" \" $1}' /tmp/out" title "time" with lines - - */ diff --git a/util-jvm/src/test/scala/com/twitter/jvm/JvmStatsTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/JvmStatsTest.scala new file mode 100644 index 0000000000..c5fcf03a4e --- /dev/null +++ b/util-jvm/src/test/scala/com/twitter/jvm/JvmStatsTest.scala @@ -0,0 +1,34 @@ +package com.twitter.jvm + +import com.twitter.finagle.stats.InMemoryStatsReceiver +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.finagle.stats.exp.MetricExpression +import org.scalatest.funsuite.AnyFunSuite + +class JvmStatsTest extends AnyFunSuite { + + private[this] val receiver = new InMemoryStatsReceiver() + JvmStats.register(receiver) + private[this] val expressions = receiver.getAllExpressionsWithLabel(ExpressionSchema.Role, "jvm") + + test("expressions are instrumented") { + expressions.foreach { + case (_, expressionSchema) => + assertMetric(expressionSchema) + } + } + + test("expressions are registered") { + assertExpressionLabels("memory_pool") + assertExpressionLabels("file_descriptors") + assertExpressionLabels("gc_cycles") + assertExpressionLabels("gc_latency") + assertExpressionLabels("jvm_uptime") + } + + private[this] def assertMetric(expressionSchema: ExpressionSchema): Unit = + assert(expressionSchema.expr.isInstanceOf[MetricExpression]) + + private[this] def assertExpressionLabels(expressionKey: String): Unit = + assert(expressions.keysIterator.exists(_.name == expressionKey)) +} diff --git a/util-jvm/src/test/scala/com/twitter/jvm/JvmTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/JvmTest.scala index 0b7102ab8d..02c1125c84 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/JvmTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/JvmTest.scala @@ -1,167 +1,157 @@ package com.twitter.jvm -import com.twitter.conversions.time._ -import com.twitter.logging.{Level, TestLogging} +import com.twitter.conversions.DurationOps._ import com.twitter.util.Time -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import java.util.logging.Level +import java.util.logging.Logger +import org.mockito.ArgumentMatchers.contains +import org.mockito.Mockito.times +import org.mockito.Mockito.verify +import org.mockito.Mockito.when +import org.scalatestplus.mockito.MockitoSugar import scala.collection.mutable +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class JvmTest extends WordSpec with TestLogging { - "Jvm" should { - class JvmHelper { - object jvm extends Jvm { - @volatile private[this] var currentSnap: Snapshot = - Snapshot(Time.epoch, Heap(0, 0, Seq()), Seq()) - - val opts = new Opts { - def compileThresh = None - } - def snapCounters = Map() - def setSnap(snap: Snapshot): Unit = { - currentSnap = snap - } - - override val executor = new MockScheduledExecutorService - - def snap = currentSnap - - def pushGc(gc: Gc): Unit = { - val gcs = snap.lastGcs filter (_.name != gc.name) - setSnap(snap.copy(lastGcs = gc +: gcs)) - } - - def forceGc() = () - def edenPool = NilJvm.edenPool - def metaspaceUsage = NilJvm.metaspaceUsage - def safepoint = NilJvm.safepoint - def applicationTime = NilJvm.applicationTime - def tenuringThreshold = NilJvm.tenuringThreshold - } +class JvmTest extends AnyFunSuite with MockitoSugar { + class JvmHelper(suppliedLogger: Option[Logger] = None) extends Jvm { + protected override def logger: Logger = suppliedLogger match { + case Some(l) => l + case None => super.logger } - "ProcessId" should { - val supported = Seq("Mac OS X", "Linux") + @volatile private[this] var currentSnap: Snapshot = + Snapshot(Time.epoch, Heap(0, 0, Seq()), Seq()) - val osIsSupported = - Option(System.getProperty("os.name")).exists(supported.contains) + val opts = new Opts { + def compileThresh = None + } + def snapCounters = Map() + def setSnap(snap: Snapshot): Unit = { + currentSnap = snap + } + + override val executor: MockScheduledExecutorService = new MockScheduledExecutorService + + def snap = currentSnap - if (osIsSupported) { - "define process id on supported platforms" in { - assert(Jvm.ProcessId.isDefined) - } - } + def pushGc(gc: Gc): Unit = { + val gcs = snap.lastGcs filter (_.name != gc.name) + setSnap(snap.copy(lastGcs = gc +: gcs)) } - "foreachGc" should { - "Capture interleaving GCs with different names" in { - val h = new JvmHelper - import h._ - val b = mutable.Buffer[Gc]() - assert(jvm.executor.schedules == List()) - jvm foreachGc { b += _ } - assert(jvm.executor.schedules.size == 1) - val Seq((r, _, _, _)) = jvm.executor.schedules - r.run() - assert(b == List()) - val gc = Gc(0, "pcopy", Time.now, 1.millisecond) - jvm.pushGc(gc) - r.run() - assert(b.size == 1) - assert(b(0) == gc) - r.run() - assert(b.size == 1) - val gc1 = gc.copy(name = "CMS") - jvm.pushGc(gc1) - r.run() - assert(b.size == 2) - assert(b(1) == gc1) - r.run() - r.run() - r.run() - assert(b.size == 2) - jvm.pushGc(gc1.copy(count = 1)) - jvm.pushGc(gc.copy(count = 1)) - r.run() - assert(b.size == 4) - assert(b(2) == gc.copy(count = 1)) - assert(b(3) == gc1.copy(count = 1)) - } - - "Complain when sampling rate is too low, every 30 minutes" in Time.withCurrentTimeFrozen { - tc => - val h = new JvmHelper - import h._ - - traceLogger(Level.DEBUG) - - jvm foreachGc { _ => /*ignore*/ - } - assert(jvm.executor.schedules.size == 1) - val Seq((r, _, _, _)) = jvm.executor.schedules - val gc = Gc(0, "pcopy", Time.now, 1.millisecond) - r.run() - jvm.pushGc(gc) - r.run() - jvm.pushGc(gc.copy(count = 2)) - r.run() - assert(logLines() == Seq("Missed 1 collections for pcopy due to sampling")) - jvm.pushGc(gc.copy(count = 10)) - assert(logLines() == Seq("Missed 1 collections for pcopy due to sampling")) - r.run() - tc.advance(29.minutes) - r.run() - assert(logLines() == Seq("Missed 1 collections for pcopy due to sampling")) - tc.advance(2.minutes) - jvm.pushGc(gc.copy(count = 12)) - r.run() - assert( - logLines() == Seq( - "Missed 1 collections for pcopy due to sampling", - "Missed 8 collections for pcopy due to sampling" - ) - ) - } + def forceGc() = () + def edenPool = NilJvm.edenPool + def metaspaceUsage = NilJvm.metaspaceUsage + def safepoint = NilJvm.safepoint + def applicationTime = NilJvm.applicationTime + def tenuringThreshold = NilJvm.tenuringThreshold + } + + test("ProcessId should define processid on supported platforms") { + val supported = Seq("Mac OS X", "Linux") + + val osIsSupported = + Option(System.getProperty("os.name")).exists(supported.contains) + + if (osIsSupported) { + assert(Jvm.ProcessId.isDefined) } + } + + test("foreachGc should capture interleaving GCs with different names") { + val jvm = new JvmHelper() + val b = mutable.Buffer[Gc]() + assert(jvm.executor.schedules == List()) + jvm foreachGc { b += _ } + assert(jvm.executor.schedules.size == 1) + val r = jvm.executor.schedules.head._1 + r.run() + assert(b == List()) + val gc = Gc(0, "pcopy", Time.now, 1.millisecond) + jvm.pushGc(gc) + r.run() + assert(b.size == 1) + assert(b(0) == gc) + r.run() + assert(b.size == 1) + val gc1 = gc.copy(name = "CMS") + jvm.pushGc(gc1) + r.run() + assert(b.size == 2) + assert(b(1) == gc1) + r.run() + r.run() + r.run() + assert(b.size == 2) + jvm.pushGc(gc1.copy(count = 1)) + jvm.pushGc(gc.copy(count = 1)) + r.run() + assert(b.size == 4) + assert(b(2) == gc.copy(count = 1)) + assert(b(3) == gc1.copy(count = 1)) + } - "monitorsGcs" should { - "queries gcs in range, in reverse chronological order" in Time.withCurrentTimeFrozen { tc => - val h = new JvmHelper - import h._ - - val query = jvm.monitorGcs(10.seconds) - assert(jvm.executor.schedules.size == 1) - val Seq((r, _, _, _)) = jvm.executor.schedules - val gc0 = Gc(0, "pcopy", Time.now, 1.millisecond) - val gc1 = Gc(1, "CMS", Time.now, 1.millisecond) - jvm.pushGc(gc1) - jvm.pushGc(gc0) - r.run() - tc.advance(9.seconds) - val gc2 = Gc(1, "pcopy", Time.now, 2.milliseconds) - jvm.pushGc(gc2) - r.run() - assert(query(10.seconds.ago) == Seq(gc2, gc1, gc0)) - tc.advance(2.seconds) - val gc3 = Gc(2, "CMS", Time.now, 1.second) - jvm.pushGc(gc3) - r.run() - assert(query(10.seconds.ago) == Seq(gc3, gc2)) - } + test("foreachGc should complain when sampling rate is too low, every 30 minutes") { + Time.withCurrentTimeFrozen { tc => + val logger = mock[Logger] + when(logger.isLoggable(Level.FINE)).thenReturn(true) + val jvm = new JvmHelper(Some(logger)) + + jvm.foreachGc(_ => () /*ignore*/ ) + assert(jvm.executor.schedules.size == 1) + val r = jvm.executor.schedules.head._1 + val gc = Gc(0, "pcopy", Time.now, 1.millisecond) + r.run() + jvm.pushGc(gc) + r.run() + jvm.pushGc(gc.copy(count = 2)) + r.run() + verify(logger, times(1)).fine(contains("Missed 1 collections for pcopy due to sampling")) + jvm.pushGc(gc.copy(count = 10)) + verify(logger, times(1)).fine(contains("Missed 1 collections for pcopy due to sampling")) + r.run() + tc.advance(29.minutes) + r.run() + verify(logger, times(1)).fine(contains("Missed 1 collections for pcopy due to sampling")) + tc.advance(2.minutes) + jvm.pushGc(gc.copy(count = 12)) + r.run() + verify(logger, times(1)).fine(contains("Missed 1 collections for pcopy due to sampling")) + verify(logger, times(1)).fine(contains("Missed 8 collections for pcopy due to sampling")) } + } - "safepoint" should { - "show an increase in total time spent after a gc" in { - val j = Jvm() - val preGc = j.safepoint - val totalTimePreGc = preGc.totalTimeMillis - System.gc() - val postGc = j.safepoint - val totalTimePostGc = postGc.totalTimeMillis - assert(totalTimePostGc > totalTimePreGc) - } + test("monitorsGcs queries gcs in range, in reverse chronological order") { + Time.withCurrentTimeFrozen { tc => + val jvm = new JvmHelper() + val query = jvm.monitorGcs(10.seconds) + assert(jvm.executor.schedules.size == 1) + val r = jvm.executor.schedules.head._1 + val gc0 = Gc(0, "pcopy", Time.now, 1.millisecond) + val gc1 = Gc(1, "CMS", Time.now, 1.millisecond) + jvm.pushGc(gc1) + jvm.pushGc(gc0) + r.run() + tc.advance(9.seconds) + val gc2 = Gc(1, "pcopy", Time.now, 2.milliseconds) + jvm.pushGc(gc2) + r.run() + assert(query(10.seconds.ago) == Seq(gc2, gc1, gc0)) + tc.advance(2.seconds) + val gc3 = Gc(2, "CMS", Time.now, 1.second) + jvm.pushGc(gc3) + r.run() + assert(query(10.seconds.ago) == Seq(gc3, gc2)) } } + + test("safepoint should show an increase in total time spent after a gc") { + val j = Jvm() + val preGc = j.safepoint + val totalTimePreGc = preGc.totalTimeMillis + System.gc() + val postGc = j.safepoint + val totalTimePostGc = postGc.totalTimeMillis + assert(totalTimePostGc > totalTimePreGc) + } } diff --git a/util-jvm/src/test/scala/com/twitter/jvm/MockScheduledExecutorService.scala b/util-jvm/src/test/scala/com/twitter/jvm/MockScheduledExecutorService.scala index 29d8bcf074..62cadac2b4 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/MockScheduledExecutorService.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/MockScheduledExecutorService.scala @@ -7,7 +7,6 @@ import java.util.concurrent.{ ScheduledFuture, TimeUnit } - import scala.collection.mutable // A mostly empty implementation so that we can successfully diff --git a/util-jvm/src/test/scala/com/twitter/jvm/NumProcsTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/NumProcsTest.scala index cffb9d0c5c..74b9ea17a4 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/NumProcsTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/NumProcsTest.scala @@ -1,8 +1,8 @@ package com.twitter.jvm -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class NumProcsTest extends FunSuite { +class NumProcsTest extends AnyFunSuite { test("return the number of available processors according to the runtime by default") { assert(System.getProperty("com.twitter.jvm.numProcs") == null) diff --git a/util-jvm/src/test/scala/com/twitter/jvm/OptsTest.scala b/util-jvm/src/test/scala/com/twitter/jvm/OptTest.scala similarity index 90% rename from util-jvm/src/test/scala/com/twitter/jvm/OptsTest.scala rename to util-jvm/src/test/scala/com/twitter/jvm/OptTest.scala index 51296886fb..c9ba8578e6 100644 --- a/util-jvm/src/test/scala/com/twitter/jvm/OptsTest.scala +++ b/util-jvm/src/test/scala/com/twitter/jvm/OptTest.scala @@ -3,12 +3,9 @@ package com.twitter.jvm import java.lang.{Boolean => JBool} import java.lang.management.ManagementFactory import javax.management.ObjectName -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class OptsTest extends FunSuite { +class OptTest extends AnyFunSuite { test("Opts") { if (System.getProperty("java.vm.name").contains("HotSpot")) { val DiagnosticName = diff --git a/util-lint/BUILD b/util-lint/BUILD index 0caeec0478..2265db4d28 100644 --- a/util-lint/BUILD +++ b/util-lint/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-lint/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-lint/src/test/scala", ], diff --git a/util-lint/src/main/scala/BUILD b/util-lint/src/main/scala/BUILD index 672435d27e..89cef8c41c 100644 --- a/util-lint/src/main/scala/BUILD +++ b/util-lint/src/main/scala/BUILD @@ -1,12 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-lint", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "util/util-core/src/main/scala", + "util/util-lint/src/main/scala/com/twitter/util/lint", ], ) diff --git a/util-lint/src/main/scala/com/twitter/util/lint/BUILD b/util-lint/src/main/scala/com/twitter/util/lint/BUILD new file mode 100644 index 0000000000..972662835b --- /dev/null +++ b/util-lint/src/main/scala/com/twitter/util/lint/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-lint", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-lint/src/main/scala/com/twitter/util/lint/Rule.scala b/util-lint/src/main/scala/com/twitter/util/lint/Rule.scala index 86ebf16ca3..006d5601e8 100644 --- a/util-lint/src/main/scala/com/twitter/util/lint/Rule.scala +++ b/util-lint/src/main/scala/com/twitter/util/lint/Rule.scala @@ -3,7 +3,7 @@ package com.twitter.util.lint import java.util.regex.Pattern /** - * A single lint rule, that when [[Rule.apply() run]] evaluates + * A single lint rule, that when [[Rule.apply run]] evaluates * whether or not there are any issues. */ trait Rule { @@ -41,13 +41,7 @@ object Rule { * * @param fn is evaluated to determine if there are any issues. */ - def apply( - category: Category, - shortName: String, - desc: String - )( - fn: => Seq[Issue] - ): Rule = { + def apply(category: Category, shortName: String, desc: String)(fn: => Seq[Issue]): Rule = { val _cat = category new Rule { def apply(): Seq[Issue] = fn diff --git a/util-lint/src/test/scala/BUILD b/util-lint/src/test/scala/BUILD index ef1a07ad6c..549eba77b8 100644 --- a/util-lint/src/test/scala/BUILD +++ b/util-lint/src/test/scala/BUILD @@ -1,10 +1,6 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-lint/src/main/scala", + "util/util-lint/src/test/scala/com/twitter/util/lint", ], ) diff --git a/util-lint/src/test/scala/com/twitter/util/lint/BUILD b/util-lint/src/test/scala/com/twitter/util/lint/BUILD new file mode 100644 index 0000000000..bfc907c64e --- /dev/null +++ b/util-lint/src/test/scala/com/twitter/util/lint/BUILD @@ -0,0 +1,16 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-lint/src/main/scala/com/twitter/util/lint", + ], +) diff --git a/util-lint/src/test/scala/com/twitter/util/lint/GlobalRulesTest.scala b/util-lint/src/test/scala/com/twitter/util/lint/GlobalRulesTest.scala index 4f048422ae..f9a0ed6105 100644 --- a/util-lint/src/test/scala/com/twitter/util/lint/GlobalRulesTest.scala +++ b/util-lint/src/test/scala/com/twitter/util/lint/GlobalRulesTest.scala @@ -1,8 +1,8 @@ package com.twitter.util.lint -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class GlobalRulesTest extends FunSuite { +class GlobalRulesTest extends AnyFunSuite { private val neverRule = Rule.apply(Category.Performance, "R2", "Good") { Nil } @@ -32,7 +32,7 @@ class GlobalRulesTest extends FunSuite { // alwaysRule should not be present assert(GlobalRules.get.iterable.size == 1) // should just be neverRule - assert(GlobalRules.get.iterable.seq.head == neverRule) + assert(GlobalRules.get.iterable.head == neverRule) } } diff --git a/util-lint/src/test/scala/com/twitter/util/lint/RuleTest.scala b/util-lint/src/test/scala/com/twitter/util/lint/RuleTest.scala index 20d4602c87..cbb9e9c43f 100644 --- a/util-lint/src/test/scala/com/twitter/util/lint/RuleTest.scala +++ b/util-lint/src/test/scala/com/twitter/util/lint/RuleTest.scala @@ -1,11 +1,8 @@ package com.twitter.util.lint -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class RuleTest extends FunSuite { +class RuleTest extends AnyFunSuite { private def withName(name: String): Rule = Rule(Category.Performance, name, "descriptive description") { Nil } diff --git a/util-lint/src/test/scala/com/twitter/util/lint/RulesTest.scala b/util-lint/src/test/scala/com/twitter/util/lint/RulesTest.scala index 780e240dd6..2d1669dcf6 100644 --- a/util-lint/src/test/scala/com/twitter/util/lint/RulesTest.scala +++ b/util-lint/src/test/scala/com/twitter/util/lint/RulesTest.scala @@ -1,11 +1,8 @@ package com.twitter.util.lint -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class RulesTest extends FunSuite { +class RulesTest extends AnyFunSuite { private var flag = false diff --git a/util-logging/BUILD b/util-logging/BUILD index 6b4c5c6451..ed171ba770 100644 --- a/util-logging/BUILD +++ b/util-logging/BUILD @@ -1,11 +1,16 @@ target( + tags = [ + "bazel-compatible", + "visibility://codestructure/packages/util-logging:legal", + ], dependencies = [ "util/util-logging/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-logging/src/test/java", "util/util-logging/src/test/scala", diff --git a/util-logging/OWNERS b/util-logging/OWNERS deleted file mode 100644 index bb5d54aad9..0000000000 --- a/util-logging/OWNERS +++ /dev/null @@ -1,6 +0,0 @@ -# hzhang is owner for ScribeHandler only -banderson -dschobel -koliver -mnakamura -roanta diff --git a/util-logging/PROJECT b/util-logging/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-logging/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-logging/src/main/scala/BUILD b/util-logging/src/main/scala/BUILD index 092bc18284..866e167235 100644 --- a/util-logging/src/main/scala/BUILD +++ b/util-logging/src/main/scala/BUILD @@ -1,17 +1,9 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-logging", - repo = artifactory, - ), - dependencies = [ - "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-stats/src/main/scala", +target( + tags = [ + "bazel-compatible", + "visibility://codestructure/packages/util-logging:legal", ], - exports = [ - "util/util-app/src/main/scala", + dependencies = [ + "util/util-logging/src/main/scala/com/twitter/logging", ], ) diff --git a/util-logging/src/main/scala/com/twitter/logging/BUILD b/util-logging/src/main/scala/com/twitter/logging/BUILD new file mode 100644 index 0000000000..46262eb7de --- /dev/null +++ b/util-logging/src/main/scala/com/twitter/logging/BUILD @@ -0,0 +1,35 @@ +# workaround for https://github.com/bazelbuild/bazel/issues/3864 +alias( + name = "logging", + target = ":internal", + tags = [ + "bazel-compatible", + "visibility://codestructure/packages/util-logging:legal", + ], +) + +scala_library( + name = "internal", + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-logging", + repo = artifactory, + ), + strict_deps = True, + tags = [ + "bazel-compatible", + ], + dependencies = [ + "util/util-app/src/main/scala", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", + ], + exports = [ + "util/util-app/src/main/scala", + ], +) diff --git a/util-logging/src/main/scala/com/twitter/logging/FileHandler.scala b/util-logging/src/main/scala/com/twitter/logging/FileHandler.scala index b7186c7fd4..75d1e8fa31 100644 --- a/util-logging/src/main/scala/com/twitter/logging/FileHandler.scala +++ b/util-logging/src/main/scala/com/twitter/logging/FileHandler.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,11 +16,21 @@ package com.twitter.logging -import java.io.{File, FileOutputStream, FilenameFilter, OutputStream} +import java.io.File +import java.io.FileOutputStream +import java.io.FilenameFilter +import java.io.OutputStream import java.nio.charset.Charset -import java.util.{Calendar, Date, logging => javalog} +import java.util.Calendar +import java.util.Date +import java.util.{logging => javalog} -import com.twitter.util.{TwitterDateFormat, HandleSignal, Return, StorageUnit, Time, Try} +import com.twitter.util.TwitterDateFormat +import com.twitter.util.HandleSignal +import com.twitter.util.Return +import com.twitter.util.StorageUnit +import com.twitter.util.Time +import com.twitter.util.Try sealed abstract class Policy object Policy { @@ -60,7 +70,7 @@ object Policy { } object FileHandler { - val UTF8 = Charset.forName("UTF-8") + val UTF8: Charset = Charset.forName("UTF-8") /** * Generates a HandlerFactory that returns a FileHandler @@ -84,7 +94,8 @@ object FileHandler { rotateCount: Int = -1, formatter: Formatter = new Formatter(), level: Option[Level] = None - ) = () => new FileHandler(filename, rollPolicy, append, rotateCount, formatter, level) + ): () => FileHandler = + () => new FileHandler(filename, rollPolicy, append, rotateCount, formatter, level) } /** @@ -97,8 +108,8 @@ class FileHandler( val append: Boolean, rotateCount: Int, formatter: Formatter, - level: Option[Level] -) extends Handler(formatter, level) { + level: Option[Level]) + extends Handler(formatter, level) { // This converts relative paths to absolute paths, as expected val (filename, name) = { @@ -183,7 +194,7 @@ class FileHandler( /** * Compute the suffix for a rolled logfile, based on the roll policy. */ - def timeSuffix(date: Date) = { + def timeSuffix(date: Date): String = { val dateFormat = rollPolicy match { case Policy.Never => TwitterDateFormat("yyyy") case Policy.SigHup => TwitterDateFormat("yyyy") @@ -223,9 +234,12 @@ class FileHandler( } case Policy.Weekly(weekday) => { next.set(Calendar.HOUR_OF_DAY, 0) - do { + + next.add(Calendar.DAY_OF_MONTH, 1) + while (next.get(Calendar.DAY_OF_WEEK) != weekday) { next.add(Calendar.DAY_OF_MONTH, 1) - } while (next.get(Calendar.DAY_OF_WEEK) != weekday) + } + Some(next) } } @@ -257,7 +271,7 @@ class FileHandler( } } - def roll() = synchronized { + def roll(): Unit = synchronized { stream.close() val newFilename = filenamePrefix + "-" + timeSuffix(new Date(openTime)) + filenameSuffix new File(filename).renameTo(new File(newFilename)) @@ -274,9 +288,7 @@ class FileHandler( if (examineRollTime) { // Only allow a single thread at a time to do a roll synchronized { - nextRollTime foreach { time => - if (Time.now.inMilliseconds > time) roll() - } + nextRollTime foreach { time => if (Time.now.inMilliseconds > time) roll() } } } diff --git a/util-logging/src/main/scala/com/twitter/logging/Formatter.scala b/util-logging/src/main/scala/com/twitter/logging/Formatter.scala index c9cb895046..725d394101 100644 --- a/util-logging/src/main/scala/com/twitter/logging/Formatter.scala +++ b/util-logging/src/main/scala/com/twitter/logging/Formatter.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -17,19 +17,22 @@ package com.twitter.logging import com.twitter.util.TwitterDateFormat -import java.text.{MessageFormat, DateFormat} +import java.text.MessageFormat +import java.text.DateFormat import java.util.regex.Pattern -import java.util.{Date, GregorianCalendar, TimeZone, logging => javalog} +import java.util.Date +import java.util.GregorianCalendar +import java.util.TimeZone +import java.util.{logging => javalog} import scala.collection.mutable +import java.{util => ju} private[logging] object Formatter { // FIXME: might be nice to unmangle some scala names here. private[logging] def formatStackTrace(t: Throwable, limit: Int): List[String] = { var out = new mutable.ListBuffer[String] if (limit > 0) { - out ++= t.getStackTrace.map { elem => - " at %s".format(elem.toString) - } + out ++= t.getStackTrace.map { elem => " at %s".format(elem.toString) } if (out.length > limit) { out.trimEnd(out.length - limit) out += " (...more...)" @@ -42,7 +45,7 @@ private[logging] object Formatter { out.toList } - val DateFormatRegex = Pattern.compile("<([^>]+)>") + val DateFormatRegex: Pattern = Pattern.compile("<([^>]+)>") val DefaultStackTraceSizeLimit = 30 val DefaultPrefix = "%.3s [] %s: " } @@ -90,8 +93,8 @@ class Formatter( val truncateAt: Int = 0, val truncateStackTracesAt: Int = Formatter.DefaultStackTraceSizeLimit, val useFullPackageNames: Boolean = false, - val prefix: String = Formatter.DefaultPrefix -) extends javalog.Formatter { + val prefix: String = Formatter.DefaultPrefix) + extends javalog.Formatter { private val matcher = Formatter.DateFormatRegex.matcher(prefix) @@ -108,7 +111,7 @@ class Formatter( /** * Calendar to use for time zone display in date-time formatting. */ - val calendar = if (timezone.isDefined) { + val calendar: ju.GregorianCalendar = if (timezone.isDefined) { new GregorianCalendar(TimeZone.getTimeZone(timezone.get)) } else { new GregorianCalendar @@ -126,14 +129,15 @@ class Formatter( /** * Return the string representation of a given log level's name */ - def formatLevelName(level: javalog.Level): String = { + def formatLevelName(level: javalog.Level, overrideFatalWithError: Boolean = false): String = { level match { case x: Level => x.name case x: javalog.Level => Logger.levels.get(x.intValue) match { case None => "%03d".format(x.intValue) - case Some(level) => level.name + case Some(level) => + if (level == Level.FATAL && overrideFatalWithError) "ERROR" else level.name } } } @@ -144,22 +148,52 @@ class Formatter( */ def lineTerminator: String = "\n" + /** + * Truncates and formats the message of the given LogRecord. Formatting entails + * inserting any LogRecord parameters into the LogRecord message string. Truncation of the + * message text is controlled by the [[truncateAt]] parameter. If it is greater than 0, + * messages will be truncated at that length and ellipses appended. Messages are then returned + * as an array of lines where each entry is the line of text between newlines. + * + * @param record a [[javalog.LogRecord]] + * @return an Array representing the formatted lines of a messages. + * @see [[javalog.LogRecord#getParameters]] + */ def formatMessageLines(record: javalog.LogRecord): Array[String] = { val message = truncateText(formatText(record)) + formatMessageLines(message, record.getThrown) + } + /** + * Assumes the given message is the processed LogRecord message (see: [[formatText(record: javalog.LogRecord)]]), + * that is, any LogRecord parameters have been properly applied to the LogRecord message string as + * well as any desired truncation or trimming. The method does not further change the incoming message, + * it only splits the message on newlines and handles formatting of any given Throwable. + * + * The passed Throwable represents the LogRecord captured exception (if any). Processes the inputs + * as described in [[formatMessageLines(record: javalog.LogRecord)]]. + * + * @param message a formatted LogRecord message + * @param thrown the LogRecord captured exception (if any). Can be null. + * @see [[javalog.LogRecord#getThrown]] + */ + private[twitter] def formatMessageLines( + message: String, + thrown: Throwable + ): Array[String] = { val containsNewLine = message.indexOf('\n') >= 0 - if (!containsNewLine && record.getThrown == null) { + if (!containsNewLine && thrown == null) { Array(message) } else { val splitOnNewlines = message.split("\n") - val numThrowLines = if (record.getThrown == null) 0 else 20 + val numThrowLines = if (thrown == null) 0 else 20 val lines = new mutable.ArrayBuffer[String](splitOnNewlines.length + numThrowLines) lines ++= splitOnNewlines - if (record.getThrown ne null) { - val traceLines = Formatter.formatStackTrace(record.getThrown, truncateStackTracesAt) - lines += record.getThrown.toString + if (thrown ne null) { + val traceLines = Formatter.formatStackTrace(thrown, truncateStackTracesAt) + lines += thrown.toString if (traceLines.nonEmpty) lines ++= traceLines } @@ -219,7 +253,7 @@ class Formatter( /** * Truncates the text from a java LogRecord, if necessary */ - def truncateText(message: String) = + def truncateText(message: String): String = if ((truncateAt > 0) && (message.length > truncateAt)) message.take(truncateAt) + "..." else @@ -230,5 +264,5 @@ class Formatter( * Formatter that logs only the text of a log message, with no prefix (no date, etc). */ object BareFormatter extends Formatter { - override def format(record: javalog.LogRecord) = formatText(record) + lineTerminator + override def format(record: javalog.LogRecord): String = formatText(record) + lineTerminator } diff --git a/util-logging/src/main/scala/com/twitter/logging/Handler.scala b/util-logging/src/main/scala/com/twitter/logging/Handler.scala index c894ca4ba8..855672ea3a 100644 --- a/util-logging/src/main/scala/com/twitter/logging/Handler.scala +++ b/util-logging/src/main/scala/com/twitter/logging/Handler.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -25,11 +25,9 @@ import java.util.{logging => javalog} abstract class Handler(val formatter: Formatter, val level: Option[Level]) extends javalog.Handler { setFormatter(formatter) - level.foreach { x => - setLevel(x) - } + level.foreach { x => setLevel(x) } - override def toString = { + override def toString: String = { "<%s level=%s formatter=%s>".format(getClass.getName, getLevel, formatter.toString) } } @@ -40,34 +38,34 @@ abstract class Handler(val formatter: Formatter, val level: Option[Level]) exten */ abstract class ProxyHandler(val handler: Handler) extends Handler(handler.formatter, handler.level) { - override def close() = handler.close() + override def close(): Unit = handler.close() - override def flush() = handler.flush() + override def flush(): Unit = handler.flush() - override def getEncoding() = handler.getEncoding() + override def getEncoding(): String = handler.getEncoding() - override def getErrorManager() = handler.getErrorManager() + override def getErrorManager(): javalog.ErrorManager = handler.getErrorManager() - override def getFilter() = handler.getFilter() + override def getFilter(): javalog.Filter = handler.getFilter() - override def getFormatter() = handler.getFormatter() + override def getFormatter(): javalog.Formatter = handler.getFormatter() - override def getLevel() = handler.getLevel() + override def getLevel(): javalog.Level = handler.getLevel() - override def isLoggable(record: javalog.LogRecord) = handler.isLoggable(record) + override def isLoggable(record: javalog.LogRecord): Boolean = handler.isLoggable(record) - override def publish(record: javalog.LogRecord) = handler.publish(record) + override def publish(record: javalog.LogRecord): Unit = handler.publish(record) - override def setEncoding(encoding: String) = handler.setEncoding(encoding) + override def setEncoding(encoding: String): Unit = handler.setEncoding(encoding) - override def setErrorManager(errorManager: javalog.ErrorManager) = + override def setErrorManager(errorManager: javalog.ErrorManager): Unit = handler.setErrorManager(errorManager) - override def setFilter(filter: javalog.Filter) = handler.setFilter(filter) + override def setFilter(filter: javalog.Filter): Unit = handler.setFilter(filter) - override def setFormatter(formatter: javalog.Formatter) = handler.setFormatter(formatter) + override def setFormatter(formatter: javalog.Formatter): Unit = handler.setFormatter(formatter) - override def setLevel(level: javalog.Level) = handler.setLevel(level) + override def setLevel(level: javalog.Level): Unit = handler.setLevel(level) } object NullHandler extends Handler(BareFormatter, None) { @@ -87,12 +85,13 @@ object StringHandler { def apply( formatter: Formatter = new Formatter(), level: Option[Level] = None - ) = () => new StringHandler(formatter, level) + ): () => StringHandler = + () => new StringHandler(formatter, level) /** * for java compatibility */ - def apply() = () => new StringHandler() + def apply(): () => StringHandler = () => new StringHandler() } /** @@ -104,19 +103,19 @@ class StringHandler(formatter: Formatter = new Formatter(), level: Option[Level] // thread-safe logging private val buffer = new StringBuffer() - def publish(record: javalog.LogRecord) = { + def publish(record: javalog.LogRecord): Unit = { buffer append getFormatter().format(record) } - def close() = {} + def close(): Unit = {} - def flush() = {} + def flush(): Unit = {} - def get = { + def get: String = { buffer.toString } - def clear() = { + def clear(): Unit = { buffer.setLength(0) buffer.trimToSize() } @@ -130,12 +129,13 @@ object ConsoleHandler { def apply( formatter: Formatter = new Formatter(), level: Option[Level] = None - ) = () => new ConsoleHandler(formatter, level) + ): () => ConsoleHandler = + () => new ConsoleHandler(formatter, level) /** * for java compatibility */ - def apply() = () => new ConsoleHandler() + def apply(): () => ConsoleHandler = () => new ConsoleHandler() } /** @@ -144,11 +144,11 @@ object ConsoleHandler { class ConsoleHandler(formatter: Formatter = new Formatter(), level: Option[Level] = None) extends Handler(formatter, level) { - def publish(record: javalog.LogRecord) = { + def publish(record: javalog.LogRecord): Unit = { System.err.print(getFormatter().format(record)) } - def close() = {} + def close(): Unit = {} - def flush() = Console.flush + def flush(): Unit = Console.flush() } diff --git a/util-logging/src/main/scala/com/twitter/logging/HasLogLevel.scala b/util-logging/src/main/scala/com/twitter/logging/HasLogLevel.scala new file mode 100644 index 0000000000..ace06737f9 --- /dev/null +++ b/util-logging/src/main/scala/com/twitter/logging/HasLogLevel.scala @@ -0,0 +1,36 @@ +package com.twitter.logging + +import scala.annotation.tailrec + +/** + * Typically mixed into `Exceptions` to indicate what [[Level]] + * they should be logged at. + * + * @see Finagle's `com.twitter.finagle.Failure`. + */ +trait HasLogLevel { + def logLevel: Level +} + +object HasLogLevel { + + /** + * Finds the first [[HasLogLevel]] for the given `Throwable` including + * its chain of causes and returns its `logLevel`. + * + * @note this finds the first [[HasLogLevel]], and should be there + * be multiple in the chain of causes, it '''will not use''' + * the most severe. + * + * @return `None` if neither `t` nor any of its causes are a + * [[HasLogLevel]]. Otherwise, returns `Some` of the first + * one found. + */ + @tailrec + def unapply(t: Throwable): Option[Level] = t match { + case null => None + case hll: HasLogLevel => Some(hll.logLevel) + case _ => unapply(t.getCause) + } + +} diff --git a/util-logging/src/main/scala/com/twitter/logging/LazyLogRecord.scala b/util-logging/src/main/scala/com/twitter/logging/LazyLogRecord.scala index 51cef970bc..ebacf29499 100644 --- a/util-logging/src/main/scala/com/twitter/logging/LazyLogRecord.scala +++ b/util-logging/src/main/scala/com/twitter/logging/LazyLogRecord.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -18,12 +18,10 @@ package com.twitter.logging import java.util.{logging => javalog} -class LazyLogRecord( - level: javalog.Level, - messageGenerator: => AnyRef -) extends LogRecord(level, "") { +class LazyLogRecord(level: javalog.Level, messageGenerator: => AnyRef) + extends LogRecord(level, "") { - override lazy val getMessage = messageGenerator.toString + override lazy val getMessage: String = messageGenerator.toString } /** @@ -32,5 +30,5 @@ class LazyLogRecord( class LazyLogRecordUnformatted(level: javalog.Level, message: String, items: Any*) extends LazyLogRecord(level, { message.format(items: _*) }) { require(items.size > 0) - val preformatted = message + val preformatted: String = message } diff --git a/util-logging/src/main/scala/com/twitter/logging/Level.scala b/util-logging/src/main/scala/com/twitter/logging/Level.scala new file mode 100644 index 0000000000..56fcfa13e4 --- /dev/null +++ b/util-logging/src/main/scala/com/twitter/logging/Level.scala @@ -0,0 +1,36 @@ +package com.twitter.logging + +import java.util.{logging => javalog} + +// replace java's ridiculous log levels with the standard ones. +sealed abstract class Level(val name: String, val value: Int) extends javalog.Level(name, value) { + // for java compat + def get(): Level = this +} + +object Level { + case object OFF extends Level("OFF", Int.MaxValue) + case object FATAL extends Level("FATAL", 1000) + case object CRITICAL extends Level("CRITICAL", 970) + case object ERROR extends Level("ERROR", 930) + case object WARNING extends Level("WARNING", 900) + case object INFO extends Level("INFO", 800) + case object DEBUG extends Level("DEBUG", 500) + case object TRACE extends Level("TRACE", 400) + case object ALL extends Level("ALL", Int.MinValue) + + private[logging] val AllLevels: Seq[Level] = + Seq(OFF, FATAL, CRITICAL, ERROR, WARNING, INFO, DEBUG, TRACE, ALL) + + /** + * Associate `java.util.logging.Level` and `Level` by their integer + * values. If there is no match, we return `None`. + */ + def fromJava(level: javalog.Level): Option[Level] = + AllLevels.find(_.value == level.intValue) + + /** + * Get a `Level` by its name, or `None` if there is no match. + */ + def parse(name: String): Option[Level] = AllLevels.find(_.name == name) +} diff --git a/util-logging/src/main/scala/com/twitter/logging/LogRecord.scala b/util-logging/src/main/scala/com/twitter/logging/LogRecord.scala index 935e0c2bc3..39a365d53e 100644 --- a/util-logging/src/main/scala/com/twitter/logging/LogRecord.scala +++ b/util-logging/src/main/scala/com/twitter/logging/LogRecord.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -19,12 +19,12 @@ package com.twitter.logging import java.util.{logging => javalog} /** - * Wrapper around {@link java.util.logging.LogRecord}. Should only be accessed from a single thread. + * Wrapper around `java.util.logging.LogRecord`. Should only be accessed from a single thread. * - * Messages are formatted by Java's LogRecord using {@link java.text.MessageFormat} whereas - * this class uses a regular {@link java.text.StringFormat} + * Messages are formatted by Java's LogRecord using `java.text.MessageFormat` whereas + * this class uses a regular `java.text.StringFormat`. * - * This class takes {@link com.twitter.logging.Logger} into account when inferring the + * This class takes [[com.twitter.logging.Logger]] into account when inferring the * `sourceMethod` and `sourceClass` names. */ class LogRecord(level: javalog.Level, msg: String) extends javalog.LogRecord(level, msg) { diff --git a/util-logging/src/main/scala/com/twitter/logging/Logger.scala b/util-logging/src/main/scala/com/twitter/logging/Logger.scala index 1916e586e7..64340ebfe4 100644 --- a/util-logging/src/main/scala/com/twitter/logging/Logger.scala +++ b/util-logging/src/main/scala/com/twitter/logging/Logger.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -18,75 +18,9 @@ package com.twitter.logging import java.util.concurrent.ConcurrentHashMap import java.util.{logging => javalog} -import scala.annotation.{tailrec, varargs} -import scala.collection.JavaConverters._ +import scala.annotation.varargs import scala.collection.Map - -// replace java's ridiculous log levels with the standard ones. -sealed abstract class Level(val name: String, val value: Int) extends javalog.Level(name, value) { - // for java compat - def get(): Level = this -} - -object Level { - case object OFF extends Level("OFF", Int.MaxValue) - case object FATAL extends Level("FATAL", 1000) - case object CRITICAL extends Level("CRITICAL", 970) - case object ERROR extends Level("ERROR", 930) - case object WARNING extends Level("WARNING", 900) - case object INFO extends Level("INFO", 800) - case object DEBUG extends Level("DEBUG", 500) - case object TRACE extends Level("TRACE", 400) - case object ALL extends Level("ALL", Int.MinValue) - - private[logging] val AllLevels: Seq[Level] = - Seq(OFF, FATAL, CRITICAL, ERROR, WARNING, INFO, DEBUG, TRACE, ALL) - - /** - * Associate [[java.util.logging.Level]] and `Level` by their integer - * values. If there is no match, we return `None`. - */ - def fromJava(level: javalog.Level): Option[Level] = - AllLevels.find(_.value == level.intValue) - - /** - * Get a `Level` by its name, or `None` if there is no match. - */ - def parse(name: String): Option[Level] = AllLevels.find(_.name == name) -} - -/** - * Typically mixed into `Exceptions` to indicate what [[Level]] - * they should be logged at. - * - * @see Finagle's `com.twitter.finagle.Failure`. - */ -trait HasLogLevel { - def logLevel: Level -} - -object HasLogLevel { - - /** - * Finds the first [[HasLogLevel]] for the given `Throwable` including - * its chain of causes and returns its `logLevel`. - * - * @note this finds the first [[HasLogLevel]], and should be there - * be multiple in the chain of causes, it '''will not use''' - * the most severe. - * - * @return `None` if neither `t` nor any of its causes are a - * [[HasLogLevel]]. Otherwise, returns `Some` of the first - * one found. - */ - @tailrec - def unapply(t: Throwable): Option[Level] = t match { - case null => None - case hll: HasLogLevel => Some(hll.logLevel) - case _ => unapply(t.getCause) - } - -} +import scala.jdk.CollectionConverters._ class LoggingException(reason: String) extends Exception(reason) @@ -95,20 +29,20 @@ class LoggingException(reason: String) extends Exception(reason) */ class Logger protected (val name: String, private val wrapped: javalog.Logger) { // wrapped methods: - def addHandler(handler: javalog.Handler) = wrapped.addHandler(handler) - def getFilter() = wrapped.getFilter() - def getHandlers() = wrapped.getHandlers() - def getLevel() = wrapped.getLevel() - def getParent() = wrapped.getParent() - def getUseParentHandlers() = wrapped.getUseParentHandlers() - def isLoggable(level: javalog.Level) = wrapped.isLoggable(level) - def log(record: javalog.LogRecord) = wrapped.log(record) - def removeHandler(handler: javalog.Handler) = wrapped.removeHandler(handler) - def setFilter(filter: javalog.Filter) = wrapped.setFilter(filter) - def setLevel(level: javalog.Level) = wrapped.setLevel(level) - def setUseParentHandlers(use: Boolean) = wrapped.setUseParentHandlers(use) - - override def toString = { + def addHandler(handler: javalog.Handler): Unit = wrapped.addHandler(handler) + def getFilter(): javalog.Filter = wrapped.getFilter() + def getHandlers(): Array[javalog.Handler] = wrapped.getHandlers() + def getLevel(): javalog.Level = wrapped.getLevel() + def getParent(): javalog.Logger = wrapped.getParent() + def getUseParentHandlers(): Boolean = wrapped.getUseParentHandlers() + def isLoggable(level: javalog.Level): Boolean = wrapped.isLoggable(level) + def log(record: javalog.LogRecord): Unit = wrapped.log(record) + def removeHandler(handler: javalog.Handler): Unit = wrapped.removeHandler(handler) + def setFilter(filter: javalog.Filter): Unit = wrapped.setFilter(filter) + def setLevel(level: javalog.Level): Unit = wrapped.setLevel(level) + def setUseParentHandlers(use: Boolean): Unit = wrapped.setUseParentHandlers(use) + + override def toString: String = { "<%s name='%s' level=%s handlers=%s use_parent=%s>".format( getClass.getName, name, @@ -132,7 +66,7 @@ class Logger protected (val name: String, private val wrapped: javalog.Logger) { */ @varargs final def log(level: Level, thrown: Throwable, message: String, items: Any*): Unit = { - val myLevel = getLevel + val myLevel = getLevel() if ((myLevel eq null) || (level.intValue >= myLevel.intValue)) { val record = @@ -146,42 +80,47 @@ class Logger protected (val name: String, private val wrapped: javalog.Logger) { } } - final def apply(level: Level, message: String, items: Any*) = log(level, message, items: _*) + final def apply(level: Level, message: String, items: Any*): Unit = log(level, message, items: _*) - final def apply(level: Level, thrown: Throwable, message: String, items: Any*) = - log(level, thrown, message, items) + final def apply(level: Level, thrown: Throwable, message: String, items: Any*): Unit = + log(level, thrown, message, items: _*) // convenience methods: @varargs - def fatal(msg: String, items: Any*) = log(Level.FATAL, msg, items: _*) + def fatal(msg: String, items: Any*): Unit = log(Level.FATAL, msg, items: _*) @varargs - def fatal(thrown: Throwable, msg: String, items: Any*) = log(Level.FATAL, thrown, msg, items: _*) + def fatal(thrown: Throwable, msg: String, items: Any*): Unit = + log(Level.FATAL, thrown, msg, items: _*) @varargs - def critical(msg: String, items: Any*) = log(Level.CRITICAL, msg, items: _*) + def critical(msg: String, items: Any*): Unit = log(Level.CRITICAL, msg, items: _*) @varargs - def critical(thrown: Throwable, msg: String, items: Any*) = + def critical(thrown: Throwable, msg: String, items: Any*): Unit = log(Level.CRITICAL, thrown, msg, items: _*) @varargs - def error(msg: String, items: Any*) = log(Level.ERROR, msg, items: _*) + def error(msg: String, items: Any*): Unit = log(Level.ERROR, msg, items: _*) @varargs - def error(thrown: Throwable, msg: String, items: Any*) = log(Level.ERROR, thrown, msg, items: _*) + def error(thrown: Throwable, msg: String, items: Any*): Unit = + log(Level.ERROR, thrown, msg, items: _*) @varargs - def warning(msg: String, items: Any*) = log(Level.WARNING, msg, items: _*) + def warning(msg: String, items: Any*): Unit = log(Level.WARNING, msg, items: _*) @varargs - def warning(thrown: Throwable, msg: String, items: Any*) = + def warning(thrown: Throwable, msg: String, items: Any*): Unit = log(Level.WARNING, thrown, msg, items: _*) @varargs - def info(msg: String, items: Any*) = log(Level.INFO, msg, items: _*) + def info(msg: String, items: Any*): Unit = log(Level.INFO, msg, items: _*) @varargs - def info(thrown: Throwable, msg: String, items: Any*) = log(Level.INFO, thrown, msg, items: _*) + def info(thrown: Throwable, msg: String, items: Any*): Unit = + log(Level.INFO, thrown, msg, items: _*) @varargs - def debug(msg: String, items: Any*) = log(Level.DEBUG, msg, items: _*) + def debug(msg: String, items: Any*): Unit = log(Level.DEBUG, msg, items: _*) @varargs - def debug(thrown: Throwable, msg: String, items: Any*) = log(Level.DEBUG, thrown, msg, items: _*) + def debug(thrown: Throwable, msg: String, items: Any*): Unit = + log(Level.DEBUG, thrown, msg, items: _*) @varargs - def trace(msg: String, items: Any*) = log(Level.TRACE, msg, items: _*) + def trace(msg: String, items: Any*): Unit = log(Level.TRACE, msg, items: _*) @varargs - def trace(thrown: Throwable, msg: String, items: Any*) = log(Level.TRACE, thrown, msg, items: _*) + def trace(thrown: Throwable, msg: String, items: Any*): Unit = + log(Level.TRACE, thrown, msg, items: _*) def debugLazy(msg: => AnyRef): Unit = logLazy(Level.DEBUG, null, msg) def traceLazy(msg: => AnyRef): Unit = logLazy(Level.TRACE, null, msg) @@ -197,7 +136,7 @@ class Logger protected (val name: String, private val wrapped: javalog.Logger) { * and attach an exception and stack trace. */ def logLazy(level: Level, thrown: Throwable, message: => AnyRef): Unit = { - val myLevel = getLevel + val myLevel = getLevel() if ((myLevel eq null) || (level.intValue >= myLevel.intValue)) { val record = new LazyLogRecord(level, message) record.setLoggerName(wrapped.getName) @@ -209,20 +148,22 @@ class Logger protected (val name: String, private val wrapped: javalog.Logger) { } // convenience methods: - def ifFatal(message: => AnyRef) = logLazy(Level.FATAL, message) - def ifFatal(thrown: Throwable, message: => AnyRef) = logLazy(Level.FATAL, thrown, message) - def ifCritical(message: => AnyRef) = logLazy(Level.CRITICAL, message) - def ifCritical(thrown: Throwable, message: => AnyRef) = logLazy(Level.CRITICAL, thrown, message) - def ifError(message: => AnyRef) = logLazy(Level.ERROR, message) - def ifError(thrown: Throwable, message: => AnyRef) = logLazy(Level.ERROR, thrown, message) - def ifWarning(message: => AnyRef) = logLazy(Level.WARNING, message) - def ifWarning(thrown: Throwable, message: => AnyRef) = logLazy(Level.WARNING, thrown, message) - def ifInfo(message: => AnyRef) = logLazy(Level.INFO, message) - def ifInfo(thrown: Throwable, message: => AnyRef) = logLazy(Level.INFO, thrown, message) - def ifDebug(message: => AnyRef) = logLazy(Level.DEBUG, message) - def ifDebug(thrown: Throwable, message: => AnyRef) = logLazy(Level.DEBUG, thrown, message) - def ifTrace(message: => AnyRef) = logLazy(Level.TRACE, message) - def ifTrace(thrown: Throwable, message: => AnyRef) = logLazy(Level.TRACE, thrown, message) + def ifFatal(message: => AnyRef): Unit = logLazy(Level.FATAL, message) + def ifFatal(thrown: Throwable, message: => AnyRef): Unit = logLazy(Level.FATAL, thrown, message) + def ifCritical(message: => AnyRef): Unit = logLazy(Level.CRITICAL, message) + def ifCritical(thrown: Throwable, message: => AnyRef): Unit = + logLazy(Level.CRITICAL, thrown, message) + def ifError(message: => AnyRef): Unit = logLazy(Level.ERROR, message) + def ifError(thrown: Throwable, message: => AnyRef): Unit = logLazy(Level.ERROR, thrown, message) + def ifWarning(message: => AnyRef): Unit = logLazy(Level.WARNING, message) + def ifWarning(thrown: Throwable, message: => AnyRef): Unit = + logLazy(Level.WARNING, thrown, message) + def ifInfo(message: => AnyRef): Unit = logLazy(Level.INFO, message) + def ifInfo(thrown: Throwable, message: => AnyRef): Unit = logLazy(Level.INFO, thrown, message) + def ifDebug(message: => AnyRef): Unit = logLazy(Level.DEBUG, message) + def ifDebug(thrown: Throwable, message: => AnyRef): Unit = logLazy(Level.DEBUG, thrown, message) + def ifTrace(message: => AnyRef): Unit = logLazy(Level.TRACE, message) + def ifTrace(thrown: Throwable, message: => AnyRef): Unit = logLazy(Level.TRACE, thrown, message) /** * Lazily logs a throwable. If the throwable contains a log level (via [[HasLogLevel]]), it @@ -244,7 +185,7 @@ class Logger protected (val name: String, private val wrapped: javalog.Logger) { /** * Remove all existing log handlers. */ - def clearHandlers() = { + def clearHandlers(): Unit = { // some custom Logger implementations may return null from getHandlers val handlers = getHandlers() if (handlers ne null) { @@ -270,14 +211,10 @@ object NullLogger object Logger extends Iterable[Logger] { private[this] val levelNamesMap: Map[String, Level] = - Level.AllLevels.map { level => - level.name -> level - }.toMap + Level.AllLevels.map { level => level.name -> level }.toMap private[this] val levelsMap: Map[Int, Level] = - Level.AllLevels.map { level => - level.value -> level - }.toMap + Level.AllLevels.map { level => level.value -> level }.toMap // A cache of scala Logger objects by name. // Using a low concurrencyLevel (2), with the assumption that we aren't ever creating too @@ -292,37 +229,37 @@ object Logger extends Iterable[Logger] { // ----- convenience methods: /** OFF is used to turn off logging entirely. */ - def OFF = Level.OFF + def OFF: Level.OFF.type = Level.OFF /** Describes an event which will cause the application to exit immediately, in failure. */ - def FATAL = Level.FATAL + def FATAL: Level.FATAL.type = Level.FATAL /** Describes an event which will cause the application to fail to work correctly, but * keep attempt to continue. The application may be unusable. */ - def CRITICAL = Level.CRITICAL + def CRITICAL: Level.CRITICAL.type = Level.CRITICAL /** Describes a user-visible error that may be transient or not affect other users. */ - def ERROR = Level.ERROR + def ERROR: Level.ERROR.type = Level.ERROR /** Describes a problem which is probably not user-visible but is notable and/or may be * an early indication of a future error. */ - def WARNING = Level.WARNING + def WARNING: Level.WARNING.type = Level.WARNING /** Describes information about the normal, functioning state of the application. */ - def INFO = Level.INFO + def INFO: Level.INFO.type = Level.INFO /** Describes information useful for general debugging, but probably not normal, * day-to-day use. */ - def DEBUG = Level.DEBUG + def DEBUG: Level.DEBUG.type = Level.DEBUG /** Describes information useful for intense debugging. */ - def TRACE = Level.TRACE + def TRACE: Level.TRACE.type = Level.TRACE /** ALL is used to log everything. */ - def ALL = Level.ALL + def ALL: Level.ALL.type = Level.ALL /** * Return a map of log level values to the corresponding Level objects. @@ -339,7 +276,7 @@ object Logger extends Iterable[Logger] { * INFO level and goes to the console (stderr). Any existing log * handlers are removed. */ - def reset() = { + def reset(): Unit = { clearHandlers() loggersCache.clear() root.addHandler(new ConsoleHandler(new Formatter(), None)) @@ -348,7 +285,7 @@ object Logger extends Iterable[Logger] { /** * Remove all existing log handlers from all existing loggers. */ - def clearHandlers() = { + def clearHandlers(): Unit = { foreach { logger => logger.clearHandlers() logger.setLevel(null) @@ -400,7 +337,7 @@ object Logger extends Iterable[Logger] { } /** An alias for `get(name)` */ - def apply(name: String) = get(name) + def apply(name: String): Logger = get(name) private def get(depth: Int): Logger = { val st = new Throwable().getStackTrace @@ -421,7 +358,7 @@ object Logger extends Iterable[Logger] { def get(): Logger = get(2) /** An alias for `get()` */ - def apply() = get(2) + def apply(): Logger = get(2) private def getForClassName(className: String) = { if (className.endsWith("$")) { @@ -437,7 +374,7 @@ object Logger extends Iterable[Logger] { def get(cls: Class[_]): Logger = getForClassName(cls.getName) /** An alias for `get(class)` */ - def apply(cls: Class[_]) = get(cls) + def apply(cls: Class[_]): Logger = get(cls) /** * Iterate the Logger objects that have been created. diff --git a/util-logging/src/main/scala/com/twitter/logging/LoggerFactory.scala b/util-logging/src/main/scala/com/twitter/logging/LoggerFactory.scala index 1426552747..b94dd484d9 100644 --- a/util-logging/src/main/scala/com/twitter/logging/LoggerFactory.scala +++ b/util-logging/src/main/scala/com/twitter/logging/LoggerFactory.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -38,8 +38,8 @@ case class LoggerFactory( node: String = "", level: Option[Level] = None, handlers: List[HandlerFactory] = Nil, - useParents: Boolean = true -) extends (() => Logger) { + useParents: Boolean = true) + extends (() => Logger) { /** * Registers new handlers and setting with the logger at `node` @@ -48,12 +48,8 @@ case class LoggerFactory( def apply(): Logger = { val logger = Logger.get(node) logger.clearHandlers() - level.foreach { x => - logger.setLevel(x) - } - handlers.foreach { h => - logger.addHandler(h()) - } + level.foreach { x => logger.setLevel(x) } + handlers.foreach { h => logger.addHandler(h()) } logger.setUseParentHandlers(useParents) logger } diff --git a/util-logging/src/main/scala/com/twitter/logging/Logging.scala b/util-logging/src/main/scala/com/twitter/logging/Logging.scala index a43d1d567f..8824d99828 100644 --- a/util-logging/src/main/scala/com/twitter/logging/Logging.scala +++ b/util-logging/src/main/scala/com/twitter/logging/Logging.scala @@ -24,7 +24,7 @@ object Logging { trait Logging { self: App => import Logging._ - lazy val log = Logger(name) + lazy val log: Logger = Logger(name) def defaultFormatter: Formatter = new Formatter() def defaultOutput: String = "/dev/stderr" diff --git a/util-logging/src/main/scala/com/twitter/logging/QueueingHandler.scala b/util-logging/src/main/scala/com/twitter/logging/QueueingHandler.scala index 1035c1a05f..2cc4e8f258 100644 --- a/util-logging/src/main/scala/com/twitter/logging/QueueingHandler.scala +++ b/util-logging/src/main/scala/com/twitter/logging/QueueingHandler.scala @@ -1,19 +1,3 @@ -/* - * Copyright 2011 Twitter, Inc. - * - * 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 com.twitter.logging import com.twitter.concurrent.{NamedPoolThreadFactory, AsyncQueue} @@ -69,7 +53,7 @@ object QueueingHandler { * [[onOverflow]] when it is at capacity. * * @param inferClassNames [[com.twitter.logging.LogRecord]] and - * [[java.util.logging.LogRecord]] both attempt to infer the class and + * `java.util.logging.LogRecord` both attempt to infer the class and * method name of the caller, but the inference needs the stack trace at * the time that the record is logged. QueueingHandler breaks the * inference because the log record is rendered out of band, so the stack @@ -144,9 +128,7 @@ class QueueingHandler(handler: Handler, val maxQueueSize: Int, inferClassNames: override def flush(): Unit = { // Publish all records in queue - queue.drain().map { records => - records.foreach(doPublish) - } + queue.drain().map { records => records.foreach(doPublish) } // Propagate flush super.flush() diff --git a/util-logging/src/main/scala/com/twitter/logging/ScribeHandler.scala b/util-logging/src/main/scala/com/twitter/logging/ScribeHandler.scala deleted file mode 100644 index 23951acc6e..0000000000 --- a/util-logging/src/main/scala/com/twitter/logging/ScribeHandler.scala +++ /dev/null @@ -1,562 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.logging - -import java.io.IOException -import java.net._ -import java.nio.{ByteBuffer, ByteOrder} -import java.nio.charset.StandardCharsets.ISO_8859_1 -import java.util.concurrent.atomic.AtomicLong -import java.util.concurrent.{ArrayBlockingQueue, LinkedBlockingQueue, TimeUnit, ThreadPoolExecutor} -import java.util.{Arrays, logging => javalog} - -import com.twitter.concurrent.NamedPoolThreadFactory -import com.twitter.conversions.string._ -import com.twitter.conversions.time._ -import com.twitter.finagle.stats.{NullStatsReceiver, StatsReceiver} -import com.twitter.util.{Duration, Time} -import scala.util.control.NonFatal - -object ScribeHandler { - private sealed trait ServerType - private case object Unknown extends ServerType - private case object Archaic extends ServerType - private case object Modern extends ServerType - - val OK = 0 - val TRY_LATER = 1 - - val DefaultHostname = "localhost" - val DefaultPort = 1463 - val DefaultCategory = "scala" - val DefaultBufferTime = 100.milliseconds - val DefaultConnectBackoff = 15.seconds - val DefaultMaxMessagesPerTransaction = 1000 - val DefaultMaxMessagesToBuffer = 10000 - val DefaultStatsReportPeriod = 5.minutes - val log = Logger.get(getClass) - - /** - * Generates a HandlerFactory that returns a ScribeHandler. - * - * @note ScribeHandler is usually used to write structured binary data. - * When used in this way, wrapping this in other handlers, such as ThrottledHandler, - * which emit plain-text messages into the log, will corrupt the resulting data. - * - * @note Java users of this library can use one of the methods in [[ScribeHandlers]], - * so they need not provide all the default arguments. - * - * @param bufferTime - * send a scribe message no more frequently than this: - * - * @param connectBackoff - * don't connect more frequently than this (when the scribe server is down): - */ - def apply( - hostname: String = DefaultHostname, - port: Int = DefaultPort, - category: String = DefaultCategory, - bufferTime: Duration = DefaultBufferTime, - connectBackoff: Duration = DefaultConnectBackoff, - maxMessagesPerTransaction: Int = DefaultMaxMessagesPerTransaction, - maxMessagesToBuffer: Int = DefaultMaxMessagesToBuffer, - formatter: Formatter = new Formatter(), - level: Option[Level] = None, - statsReceiver: StatsReceiver = NullStatsReceiver - ): () => ScribeHandler = - () => - new ScribeHandler( - hostname, - port, - category, - bufferTime, - connectBackoff, - maxMessagesPerTransaction, - maxMessagesToBuffer, - formatter, - level, - statsReceiver - ) - - def apply( - hostname: String, - port: Int, - category: String, - bufferTime: Duration, - connectBackoff: Duration, - maxMessagesPerTransaction: Int, - maxMessagesToBuffer: Int, - formatter: Formatter, - level: Option[Level] - ): () => ScribeHandler = - apply( - hostname, - port, - category, - bufferTime, - connectBackoff, - maxMessagesPerTransaction, - maxMessagesToBuffer, - formatter, - level, - NullStatsReceiver - ) -} - -/** - * NOTE: ScribeHandler is usually used to write structured binary data. - * When used in this way, wrapping this in other handlers, such as ThrottledHandler, - * which emit plain-text messages into the log, will corrupt the resulting data. - */ -class ScribeHandler( - hostname: String, - port: Int, - category: String, - bufferTime: Duration, - connectBackoff: Duration, - maxMessagesPerTransaction: Int, - maxMessagesToBuffer: Int, - formatter: Formatter, - level: Option[Level], - statsReceiver: StatsReceiver -) extends Handler(formatter, level) { - import ScribeHandler._ - - def this( - hostname: String, - port: Int, - category: String, - bufferTime: Duration, - connectBackoff: Duration, - maxMessagesPerTransaction: Int, - maxMessagesToBuffer: Int, - formatter: Formatter, - level: Option[Level] - ) = - this( - hostname, - port, - category, - bufferTime, - connectBackoff, - maxMessagesPerTransaction, - maxMessagesToBuffer, - formatter, - level, - NullStatsReceiver - ) - - private[this] val stats = new ScribeHandlerStats(statsReceiver) - - // it may be necessary to log errors here if scribe is down: - private val loggerName = getClass.getName - - private var lastConnectAttempt = Time.epoch - - @volatile private var _lastTransmission = Time.epoch - // visible for testing - private[logging] def updateLastTransmission(): Unit = _lastTransmission = Time.now - - private var socket: Option[Socket] = None - // visible for testing - private[logging] def setSocket(newSocket: Option[Socket]) = synchronized { socket = newSocket } - - private var serverType: ServerType = Unknown - - private def isArchaicServer() = { serverType == Archaic } - - // Could be rewritten using a simple Condition (await/notify) or producer/consumer - // with timed batching - private[logging] val flusher = { - val threadFactory = new NamedPoolThreadFactory("ScribeFlusher-" + category, true) - // should be 1, but this is a crude form of retry - val queue = new ArrayBlockingQueue[Runnable](5) - val rejectionHandler = new ThreadPoolExecutor.DiscardPolicy() - new ThreadPoolExecutor(1, 1, 0L, TimeUnit.MILLISECONDS, queue, threadFactory, rejectionHandler) - } - - private[logging] val queue = new LinkedBlockingQueue[Array[Byte]](maxMessagesToBuffer) - - def queueSize: Int = queue.size() - - override def flush(): Unit = { - - def connect(): Unit = { - if (!socket.isDefined) { - if (Time.now.since(lastConnectAttempt) > connectBackoff) { - try { - lastConnectAttempt = Time.now - socket = Some(new Socket(hostname, port)) - - stats.incrConnection() - serverType = Unknown - } catch { - case e: Exception => - log.error("Unable to open socket to scribe server at %s:%d: %s", hostname, port, e) - stats.incrConnectionFailure() - } - } else { - stats.incrConnectionSkipped() - } - } - } - - def detectArchaicServer(): Unit = { - if (serverType != Unknown) return - - serverType = socket match { - case None => Unknown - case Some(s) => { - val outStream = s.getOutputStream() - - try { - val fakeMessageWithOldScribePrefix: Array[Byte] = { - val prefix = OLD_SCRIBE_PREFIX - val messageSize = prefix.length + 5 - val buffer = ByteBuffer.wrap(new Array[Byte](messageSize + 4)) - buffer.order(ByteOrder.BIG_ENDIAN) - buffer.putInt(messageSize) - buffer.put(prefix) - buffer.putInt(0) - buffer.put(0: Byte) - buffer.array - } - - outStream.write(fakeMessageWithOldScribePrefix) - readResponseExpecting(s, OLD_SCRIBE_REPLY) - - // Didn't get exception, so the server must be archaic. - log.debug("Scribe server is archaic; changing to old protocol for future requests.") - Archaic - } catch { - case NonFatal(_) => Modern - } - } - } - } - - // TODO we should send any remaining messages before the app terminates - def sendBatch(): Unit = { - synchronized { - connect() - detectArchaicServer() - socket.foreach { s => - val outStream = s.getOutputStream() - var remaining = queue.size - // try to send the log records in batch - // need to check if socket is closed due to exception - while (remaining > 0 && socket.isDefined) { - val count = maxMessagesPerTransaction min remaining - val buffer = makeBuffer(count) - var offset = 0 - - try { - outStream.write(buffer.array) - val expectedReply = if (isArchaicServer()) OLD_SCRIBE_REPLY else SCRIBE_REPLY - - readResponseExpecting(s, expectedReply) - - stats.incrSentRecords(count) - remaining -= count - } catch { - case e: Exception => - stats.incrDroppedRecords(count) - log.error( - e, - "Failed to send %s %d log entries to scribe server at %s:%d", - category, - count, - hostname, - port - ) - closeSocket() - } - } - updateLastTransmission() - } - stats.log() - } - } - - def readResponseExpecting(socket: Socket, expectedReply: Array[Byte]): Unit = { - var offset = 0 - - val inStream = socket.getInputStream() - - // read response: - val response = new Array[Byte](expectedReply.length) - while (offset < response.length) { - val n = inStream.read(response, offset, response.length - offset) - if (n < 0) { - throw new IOException("End of stream") - } - offset += n - } - if (!Arrays.equals(response, expectedReply)) { - throw new IOException("Error response from scribe server: " + response.hexlify) - } - } - - flusher.execute(new Runnable { - def run(): Unit = { sendBatch() } - }) - } - - // should be private, make it visible to tests - private[logging] def makeBuffer(count: Int): ByteBuffer = { - val texts = for (i <- 0 until count) yield queue.poll() - - val recordHeader = ByteBuffer.wrap(new Array[Byte](10 + category.length)) - recordHeader.order(ByteOrder.BIG_ENDIAN) - recordHeader.put(11: Byte) - recordHeader.putShort(1) - recordHeader.putInt(category.length) - recordHeader.put(category.getBytes(ISO_8859_1)) - recordHeader.put(11: Byte) - recordHeader.putShort(2) - - val prefix = if (isArchaicServer()) OLD_SCRIBE_PREFIX else SCRIBE_PREFIX - val messageSize = (count * (recordHeader.capacity + 5)) + texts.foldLeft(0) { _ + _.length } + prefix.length + 5 - val buffer = ByteBuffer.wrap(new Array[Byte](messageSize + 4)) - buffer.order(ByteOrder.BIG_ENDIAN) - // "framing": - buffer.putInt(messageSize) - buffer.put(prefix) - buffer.putInt(count) - for (text <- texts) { - buffer.put(recordHeader.array) - buffer.putInt(text.length) - buffer.put(text) - buffer.put(0: Byte) - } - buffer.put(0: Byte) - buffer - } - - private def closeSocket(): Unit = { - synchronized { - socket.foreach { s => - try { - s.close() - } catch { - case _: Throwable => - } - } - socket = None - } - } - - override def close(): Unit = { - stats.incrCloses() - // TODO consider draining the flusher queue before returning - // nothing stops a pending flush from opening a new socket - closeSocket() - flusher.shutdown() - } - - override def publish(record: javalog.LogRecord): Unit = { - if (record.getLoggerName == loggerName) return - publish(getFormatter.format(record).getBytes("UTF-8")) - } - - def publish(record: Array[Byte]): Unit = { - stats.incrPublished() - if (!queue.offer(record)) stats.incrDroppedRecords() - if (Time.now.since(_lastTransmission) >= bufferTime) flush() - } - - override def toString = { - ("<%s level=%s hostname=%s port=%d scribe_buffer=%s " + - "scribe_backoff=%s scribe_max_packet_size=%d formatter=%s>").format( - getClass.getName, - getLevel, - hostname, - port, - bufferTime, - connectBackoff, - maxMessagesPerTransaction, - formatter.toString - ) - } - - private[this] val SCRIBE_PREFIX: Array[Byte] = Array[Byte]( - // version 1, call, "Log", reqid=0 - 0x80.toByte, - 1, - 0, - 1, - 0, - 0, - 0, - 3, - 'L'.toByte, - 'o'.toByte, - 'g'.toByte, - 0, - 0, - 0, - 0, - // list of structs - 15, - 0, - 1, - 12 - ) - private[this] val OLD_SCRIBE_PREFIX: Array[Byte] = Array[Byte]( - // (no version), "Log", reply, reqid=0 - 0, - 0, - 0, - 3, - 'L'.toByte, - 'o'.toByte, - 'g'.toByte, - 1, - 0, - 0, - 0, - 0, - // list of structs - 15, - 0, - 1, - 12 - ) - - private[this] val SCRIBE_REPLY: Array[Byte] = Array[Byte]( - // version 1, reply, "Log", reqid=0 - 0x80.toByte, - 1, - 0, - 2, - 0, - 0, - 0, - 3, - 'L'.toByte, - 'o'.toByte, - 'g'.toByte, - 0, - 0, - 0, - 0, - // int, fid 0, 0=ok - 8, - 0, - 0, - 0, - 0, - 0, - 0, - 0 - ) - private[this] val OLD_SCRIBE_REPLY: Array[Byte] = Array[Byte]( - 0, - 0, - 0, - 20, - // (no version), "Log", reply, reqid=0 - 0, - 0, - 0, - 3, - 'L'.toByte, - 'o'.toByte, - 'g'.toByte, - 2, - 0, - 0, - 0, - 0, - // int, fid 0, 0=ok - 8, - 0, - 0, - 0, - 0, - 0, - 0, - 0 - ) - - private class ScribeHandlerStats(statsReceiver: StatsReceiver) { - private var _lastLogStats = Time.epoch - - private val sentRecords = new AtomicLong() - private val droppedRecords = new AtomicLong() - private val connectionFailure = new AtomicLong() - private val connectionSkipped = new AtomicLong() - - val totalSentRecords = statsReceiver.counter("sent_records") - val totalDroppedRecords = statsReceiver.counter("dropped_records") - val totalConnectionFailure = statsReceiver.counter("connection_failed") - val totalConnectionSkipped = statsReceiver.counter("connection_skipped") - val totalConnects = statsReceiver.counter("connects") - val totalPublished = statsReceiver.counter("published") - val totalCloses = statsReceiver.counter("closes") - val instances = statsReceiver.addGauge("instances") { 1 } - val unsentQueue = statsReceiver.addGauge("unsent_queue") { queueSize } - - def incrSentRecords(count: Int): Unit = { - sentRecords.addAndGet(count) - totalSentRecords.incr(count) - } - - def incrDroppedRecords(count: Int = 1): Unit = { - droppedRecords.addAndGet(count) - totalDroppedRecords.incr(count) - } - - def incrConnectionFailure(): Unit = { - connectionFailure.incrementAndGet() - totalConnectionFailure.incr() - } - - def incrConnectionSkipped(): Unit = { - connectionSkipped.incrementAndGet() - totalConnectionSkipped.incr() - } - - def incrPublished(): Unit = totalPublished.incr() - - def incrConnection(): Unit = totalConnects.incr() - - def incrCloses(): Unit = totalCloses.incr() - - def log(): Unit = { - synchronized { - val period = Time.now.since(_lastLogStats) - if (period > ScribeHandler.DefaultStatsReportPeriod) { - val sent = sentRecords.getAndSet(0) - val dropped = droppedRecords.getAndSet(0) - val failed = connectionFailure.getAndSet(0) - val skipped = connectionSkipped.getAndSet(0) - ScribeHandler.log.debug( - "sent records: %d, per second: %d, dropped records: %d, reconnection failures: %d, reconnection skipped: %d", - sent, - sent / period.inSeconds, - dropped, - failed, - skipped - ) - - _lastLogStats = Time.now - } - } - } - } -} diff --git a/util-logging/src/main/scala/com/twitter/logging/ScribeHandlers.scala b/util-logging/src/main/scala/com/twitter/logging/ScribeHandlers.scala deleted file mode 100644 index 5e42b2a764..0000000000 --- a/util-logging/src/main/scala/com/twitter/logging/ScribeHandlers.scala +++ /dev/null @@ -1,28 +0,0 @@ -package com.twitter.logging - -import com.twitter.finagle.stats.NullStatsReceiver -import com.twitter.logging.ScribeHandler._ - -object ScribeHandlers { - /** - * For Java compatibility. Calls [[ScribeHandler.apply()]] method with - * the provided arguments, so Java users don't need to provide the - * other default arguments in the call. - */ - def apply( - category: String, - formatter: Formatter - ): () => ScribeHandler = - ScribeHandler.apply( - hostname = DefaultHostname, - port = DefaultPort, - category = category, - bufferTime = DefaultBufferTime, - connectBackoff = DefaultConnectBackoff, - maxMessagesPerTransaction = DefaultMaxMessagesPerTransaction, - maxMessagesToBuffer = DefaultMaxMessagesToBuffer, - formatter = formatter, - level = None, - statsReceiver = NullStatsReceiver - ) -} diff --git a/util-logging/src/main/scala/com/twitter/logging/SyslogHandler.scala b/util-logging/src/main/scala/com/twitter/logging/SyslogHandler.scala index 016495d7cb..5278e589bf 100644 --- a/util-logging/src/main/scala/com/twitter/logging/SyslogHandler.scala +++ b/util-logging/src/main/scala/com/twitter/logging/SyslogHandler.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -22,6 +22,8 @@ import java.util.{logging => javalog} import com.twitter.concurrent.NamedPoolThreadFactory import com.twitter.util.{TwitterDateFormat, NetUtil} +import java.text.{DateFormat, SimpleDateFormat} +import java.util.concurrent.Future object SyslogHandler { val DEFAULT_PORT = 514 @@ -65,8 +67,8 @@ object SyslogHandler { } } - val ISO_DATE_FORMAT = TwitterDateFormat("yyyy-MM-dd'T'HH:mm:ss") - val OLD_SYSLOG_DATE_FORMAT = TwitterDateFormat("MMM dd HH:mm:ss") + val ISO_DATE_FORMAT: SimpleDateFormat = TwitterDateFormat("yyyy-MM-dd'T'HH:mm:ss") + val OLD_SYSLOG_DATE_FORMAT: SimpleDateFormat = TwitterDateFormat("MMM dd HH:mm:ss") /** * Generates a HandlerFactory that returns a SyslogHandler. @@ -82,7 +84,7 @@ object SyslogHandler { port: Int = SyslogHandler.DEFAULT_PORT, formatter: Formatter = new Formatter(), level: Option[Level] = None - ) = () => new SyslogHandler(server, port, formatter, level) + ): () => SyslogHandler = () => new SyslogHandler(server, port, formatter, level) } class SyslogHandler(val server: String, val port: Int, formatter: Formatter, level: Option[Level]) @@ -91,10 +93,10 @@ class SyslogHandler(val server: String, val port: Int, formatter: Formatter, lev private val socket = new DatagramSocket private[logging] val dest = new InetSocketAddress(server, port) - def flush() = {} - def close() = {} + def flush(): Unit = {} + def close(): Unit = {} - def publish(record: javalog.LogRecord) = { + def publish(record: javalog.LogRecord): Unit = { val data = formatter.format(record).getBytes val packet = new DatagramPacket(data, data.length, dest) SyslogFuture { @@ -138,8 +140,8 @@ class SyslogFormatter( val priority: Int = SyslogHandler.PRIORITY_USER, timezone: Option[String] = None, truncateAt: Int = 0, - truncateStackTracesAt: Int = Formatter.DefaultStackTraceSizeLimit -) extends Formatter( + truncateStackTracesAt: Int = Formatter.DefaultStackTraceSizeLimit) + extends Formatter( timezone, truncateAt, truncateStackTracesAt, @@ -147,7 +149,7 @@ class SyslogFormatter( prefix = "" ) { - override def dateFormat = + override def dateFormat: DateFormat = if (useIsoDateFormat) { SyslogHandler.ISO_DATE_FORMAT } else { @@ -176,7 +178,7 @@ object SyslogFuture { ) private val noop = new Runnable { def run(): Unit = {} } - def apply(action: => Unit) = + def apply(action: => Unit): Future[_] = executor.submit(new Runnable { def run(): Unit = { action } }) diff --git a/util-logging/src/main/scala/com/twitter/logging/ThrottledHandler.scala b/util-logging/src/main/scala/com/twitter/logging/ThrottledHandler.scala index f7b237ba32..9108d7970e 100644 --- a/util-logging/src/main/scala/com/twitter/logging/ThrottledHandler.scala +++ b/util-logging/src/main/scala/com/twitter/logging/ThrottledHandler.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,14 +16,14 @@ package com.twitter.logging -import java.util.concurrent.atomic.AtomicReference +import com.twitter.conversions.DurationOps._ +import com.twitter.util.{Duration, Time} +import java.util.concurrent.ConcurrentHashMap +import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} import java.util.{logging => javalog} - +import java.util.function.{Function => JFunction} import scala.annotation.tailrec -import scala.collection.mutable - -import com.twitter.conversions.time._ -import com.twitter.util.{Duration, Time} +import scala.jdk.CollectionConverters._ object ThrottledHandler { @@ -32,7 +32,6 @@ object ThrottledHandler { * NOTE: ThrottledHandler emits plain-text messages regarding any throttling it does. * This means that using it to wrap any logger which you expect to produce easily parseable, * well-structured logs (as opposed to just plain text logs) will break your format. - * Specifically, wrapping ScribeHandler with ThrottledHandler is usually a bug. * * @param handler * Wrapped handler. @@ -48,7 +47,9 @@ object ThrottledHandler { handler: HandlerFactory, duration: Duration = 0.seconds, maxToDisplay: Int = Int.MaxValue - ) = () => new ThrottledHandler(handler(), duration, maxToDisplay) + ): () => ThrottledHandler = () => new ThrottledHandler(handler(), duration, maxToDisplay) + + private val OneSecond = 1.second } /** @@ -66,39 +67,35 @@ object ThrottledHandler { * @param maxToDisplay * Maximum duplicate log entries to pass before suppressing them. */ -class ThrottledHandler( - handler: Handler, - val duration: Duration, - val maxToDisplay: Int -) extends ProxyHandler(handler) { +class ThrottledHandler(handler: Handler, val duration: Duration, val maxToDisplay: Int) + extends ProxyHandler(handler) { + // we accept some raciness here. it means that sometimes extra log lines will + // get logged. this was deemed preferable to holding locks. private class Throttle(startTime: Time, name: String, level: javalog.Level) { + @volatile private[this] var expired = false - private[this] var count = 0 - override def toString = "Throttle: startTime=" + startTime + " count=" + count + private[this] val count = new AtomicInteger(0) + + override def toString: String = "Throttle: startTime=" + startTime + " count=" + count.get() final def add(record: javalog.LogRecord, now: Time): Boolean = { - val (shouldPublish, added) = synchronized { - if (!expired) { - count += 1 - (count <= maxToDisplay, true) - } else { - (false, false) + if (!expired) { + if (count.incrementAndGet() <= maxToDisplay) { + doPublish(record) } + true + } else { + false } - - if (shouldPublish) doPublish(record) - added } final def removeIfExpired(now: Time): Boolean = { - val didExpire = synchronized { - expired = (now - startTime >= duration) - expired - } + val didExpire = now - startTime >= duration + expired = didExpire - if (didExpire && count > maxToDisplay) publishSwallowed() + if (didExpire && count.get() > maxToDisplay) publishSwallowed() didExpire } @@ -106,7 +103,7 @@ class ThrottledHandler( private[this] def publishSwallowed(): Unit = { val throttledRecord = new javalog.LogRecord( level, - "(swallowed %d repeating messages)".format(count - maxToDisplay) + "(swallowed %d repeating messages)".format(count.get() - maxToDisplay) ) throttledRecord.setLoggerName(name) doPublish(throttledRecord) @@ -114,7 +111,7 @@ class ThrottledHandler( } private val lastFlushCheck = new AtomicReference(Time.epoch) - private val throttleMap = new mutable.HashMap[String, Throttle] + private val throttleMap = new ConcurrentHashMap[String, Throttle]() @deprecated("Use flushThrottled() instead", "5.3.13") def reset(): Unit = { @@ -125,11 +122,9 @@ class ThrottledHandler( * Force printing any "swallowed" messages. */ def flushThrottled(): Unit = { - synchronized { - val now = Time.now - throttleMap retain { - case (_, throttle) => !throttle.removeIfExpired(now) - } + val now = Time.now + throttleMap.asScala.retain { + case (_, throttle) => !throttle.removeIfExpired(now) } } @@ -141,7 +136,7 @@ class ThrottledHandler( val now = Time.now val last = lastFlushCheck.get - if (now - last > 1.second && lastFlushCheck.compareAndSet(last, now)) { + if (now - last > ThrottledHandler.OneSecond && lastFlushCheck.compareAndSet(last, now)) { flushThrottled() } @@ -150,12 +145,13 @@ class ThrottledHandler( case _ => record.getMessage } @tailrec def tryPublish(): Unit = { - val throttle = synchronized { - throttleMap.getOrElseUpdate( - key, - new Throttle(now, record.getLoggerName(), record.getLevel()) - ) - } + val throttle = throttleMap.computeIfAbsent( + key, + new JFunction[String, Throttle] { + def apply(key: String): Throttle = + new Throttle(now, record.getLoggerName(), record.getLevel()) + } + ) // catch the case where throttle is removed before we had a chance to add if (!throttle.add(record, now)) tryPublish() @@ -164,7 +160,7 @@ class ThrottledHandler( tryPublish() } - private def doPublish(record: javalog.LogRecord) = { + private def doPublish(record: javalog.LogRecord): Unit = { super.publish(record) } } diff --git a/util-logging/src/main/scala/com/twitter/logging/config/LoggerConfig.scala b/util-logging/src/main/scala/com/twitter/logging/config/LoggerConfig.scala deleted file mode 100644 index 509d6bafbd..0000000000 --- a/util-logging/src/main/scala/com/twitter/logging/config/LoggerConfig.scala +++ /dev/null @@ -1,269 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.logging -package config - -import com.twitter.conversions.time._ -import com.twitter.util.{Config, Duration, NetUtil} - -@deprecated("use LoggerFactory", "6.12.1") -class LoggerConfig extends Config[Logger] { - - /** - * Name of the logging node. The default ("") is the top-level logger. - */ - var node: String = "" - - /** - * Log level for this node. Leaving it None is java's secret signal to use the parent logger's - * level. - */ - var level: Option[Level] = None - - /** - * Where to send log messages. - */ - var handlers: List[HandlerConfig] = Nil - - /** - * Override to have log messages stop at this node. Otherwise they are passed up to parent - * nodes. - */ - var useParents = true - - def apply(): Logger = { - val logger = Logger.get(node) - level.foreach { x => - logger.setLevel(x) - } - handlers.foreach { h => - logger.addHandler(h()) - } - logger.setUseParentHandlers(useParents) - logger - } -} - -@deprecated("use Formatter directly", "6.12.1") -class FormatterConfig extends Config[Formatter] { - - /** - * Should dates in log messages be reported in a different time zone rather than local time? - * If set, the time zone name must be one known by the java `TimeZone` class. - */ - var timezone: Option[String] = None - - /** - * Truncate log messages after N characters. 0 = don't truncate (the default). - */ - var truncateAt: Int = 0 - - /** - * Truncate stack traces in exception logging (line count). - */ - var truncateStackTracesAt: Int = 30 - - /** - * Use full package names like "com.example.thingy" instead of just the toplevel name like - * "thingy"? - */ - var useFullPackageNames: Boolean = false - - /** - * Format for the log-line prefix, if any. - * - * There are two positional format strings (printf-style): the name of the level being logged - * (for example, "ERROR") and the name of the package that's logging (for example, "jobs"). - * - * A string in `<` angle brackets `>` will be used to format the log entry's timestamp, using - * java's `SimpleDateFormat`. - * - * For example, a format string of: - * - * "%.3s [] %s: " - * - * will generate a log line prefix of: - * - * "ERR [20080315-18:39:05.033] jobs: " - */ - var prefix: String = "%.3s [] %s: " - - def apply() = - new Formatter(timezone, truncateAt, truncateStackTracesAt, useFullPackageNames, prefix) -} - -@deprecated("use BareFormatter directly", "6.12.1") -object BareFormatterConfig extends FormatterConfig { - override def apply() = BareFormatter -} - -@deprecated("use SyslogFormatter directly", "6.12.1") -class SyslogFormatterConfig extends FormatterConfig { - - /** - * Hostname to prepend to log lines. - */ - var hostname: String = NetUtil.getLocalHostName() - - /** - * Optional server name to insert before log entries. - */ - var serverName: Option[String] = None - - /** - * Use new standard ISO-format timestamps instead of old BSD-format? - */ - var useIsoDateFormat: Boolean = true - - /** - * Priority level in syslog numbers. - */ - var priority: Int = SyslogHandler.PRIORITY_USER - - def serverName_=(name: String): Unit = { serverName = Some(name) } - - override def apply() = - new SyslogFormatter( - hostname, - serverName, - useIsoDateFormat, - priority, - timezone, - truncateAt, - truncateStackTracesAt - ) -} - -@deprecated("use Formatter directly", "6.12.1") -trait HandlerConfig extends Config[Handler] { - var formatter: FormatterConfig = new FormatterConfig - - var level: Option[Level] = None -} - -@deprecated("use HandlerFactory", "6.12.1") -class ConsoleHandlerConfig extends HandlerConfig { - def apply() = new ConsoleHandler(formatter(), level) -} - -@deprecated("use HandlerFactory", "6.12.1") -class ThrottledHandlerConfig extends HandlerConfig { - - /** - * Timespan to consider duplicates. After this amount of time, duplicate entries will be logged - * again. - */ - var duration: Duration = 0.seconds - - /** - * Maximum duplicate log entries to pass before suppressing them. - */ - var maxToDisplay: Int = Int.MaxValue - - /** - * Wrapped handler. - */ - var handler: HandlerConfig = null - - def apply() = new ThrottledHandler(handler(), duration, maxToDisplay) -} - -@deprecated("use HandlerFactory", "6.12.1") -class QueuingHandlerConfig extends HandlerConfig { - def apply() = new QueueingHandler(handler(), maxQueueSize) - - /** - * Maximum queue size. Records are dropped when queue overflows. - */ - var maxQueueSize: Int = Int.MaxValue - - /** - * Wrapped handler. - */ - var handler: HandlerConfig = null -} - -@deprecated("use HandlerFactory", "6.12.1") -class FileHandlerConfig extends HandlerConfig { - - /** - * Filename to log to. - */ - var filename: String = null - - /** - * When to roll the logfile. - */ - var roll: Policy = Policy.Never - - /** - * Append to an existing logfile, or truncate it? - */ - var append: Boolean = true - - /** - * How many rotated logfiles to keep around, maximum. -1 means to keep them all. - */ - var rotateCount: Int = -1 - - def apply() = new FileHandler(filename, roll, append, rotateCount, formatter(), level) -} - -@deprecated("use HandlerFactory", "6.12.1") -class SyslogHandlerConfig extends HandlerConfig { - - /** - * Syslog server hostname. - */ - var server: String = "localhost" - - /** - * Syslog server port. - */ - var port: Int = SyslogHandler.DEFAULT_PORT - - def apply() = new SyslogHandler(server, port, formatter(), level) -} - -@deprecated("use HandlerFactory", "6.12.1") -class ScribeHandlerConfig extends HandlerConfig { - // send a scribe message no more frequently than this: - var bufferTime = 100.milliseconds - - // don't connect more frequently than this (when the scribe server is down): - var connectBackoff = 15.seconds - - var maxMessagesPerTransaction = 1000 - var maxMessagesToBuffer = 10000 - - var hostname = "localhost" - var port = 1463 - var category = "scala" - - def apply() = - new ScribeHandler( - hostname, - port, - category, - bufferTime, - connectBackoff, - maxMessagesPerTransaction, - maxMessagesToBuffer, - formatter(), - level - ) -} diff --git a/util-logging/src/main/scala/com/twitter/logging/config/package.scala b/util-logging/src/main/scala/com/twitter/logging/config/package.scala deleted file mode 100644 index 0d14da5e81..0000000000 --- a/util-logging/src/main/scala/com/twitter/logging/config/package.scala +++ /dev/null @@ -1,24 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.logging - -package object config { - type Level = com.twitter.logging.Level - val Level = com.twitter.logging.Level - type Policy = com.twitter.logging.Policy - val Policy = com.twitter.logging.Policy -} diff --git a/util-logging/src/main/scala/com/twitter/logging/package.scala b/util-logging/src/main/scala/com/twitter/logging/package.scala index a9df0daab2..80d214e09b 100644 --- a/util-logging/src/main/scala/com/twitter/logging/package.scala +++ b/util-logging/src/main/scala/com/twitter/logging/package.scala @@ -5,7 +5,7 @@ * 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 + * https://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, diff --git a/util-logging/src/test/java/BUILD b/util-logging/src/test/java/BUILD index 2a9bf3bb10..1d4116e47d 100644 --- a/util-logging/src/test/java/BUILD +++ b/util-logging/src/test/java/BUILD @@ -1,8 +1,15 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "util/util-logging/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-logging/src/main/scala/com/twitter/logging", ], ) diff --git a/util-logging/src/test/java/com/twitter/logging/LoggerFactoryTest.java b/util-logging/src/test/java/com/twitter/logging/LoggerFactoryTest.java index 8d39469eb6..345e483764 100644 --- a/util-logging/src/test/java/com/twitter/logging/LoggerFactoryTest.java +++ b/util-logging/src/test/java/com/twitter/logging/LoggerFactoryTest.java @@ -1,15 +1,20 @@ package com.twitter.logging; -import junit.framework.TestCase; +import org.junit.After; +import org.junit.Before; +import org.junit.Test; +import static org.junit.Assert.*; import scala.Function0; // Just make sure this compiles -public class LoggerFactoryTest extends TestCase { +public class LoggerFactoryTest { + @Test public void testBuilder() { LoggerFactoryBuilder builder = LoggerFactory.newBuilder(); } + @Test public void testFactory() { LoggerFactoryBuilder builder = LoggerFactory.newBuilder(); @@ -24,6 +29,7 @@ public void testFactory() { LoggerFactory factory = intermediate.build(); } + @Test public void testAddHandler() { LoggerFactoryBuilder builder = LoggerFactory.newBuilder(); diff --git a/util-logging/src/test/java/com/twitter/logging/ScribeHandlersCompilationTest.java b/util-logging/src/test/java/com/twitter/logging/ScribeHandlersCompilationTest.java deleted file mode 100644 index bb316b84cc..0000000000 --- a/util-logging/src/test/java/com/twitter/logging/ScribeHandlersCompilationTest.java +++ /dev/null @@ -1,14 +0,0 @@ -package com.twitter.logging; - -import scala.Function0; - -/** - * Basic compilation test for calling {{{@link ScribeHandlers}}} - * methods from java - */ -public class ScribeHandlersCompilationTest { - public void testScribeHandlers() { - Function0 scribeHandlerFunc = - ScribeHandlers.apply("category", BareFormatter$.MODULE$); - } -} diff --git a/util-logging/src/test/scala/BUILD b/util-logging/src/test/scala/BUILD index 949edcdc96..07b37c9ff5 100644 --- a/util-logging/src/test/scala/BUILD +++ b/util-logging/src/test/scala/BUILD @@ -1,13 +1,20 @@ junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-logging/src/main/scala", - "util/util-stats/src/main/scala", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-logging/src/main/scala/com/twitter/logging", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", ], ) diff --git a/util-logging/src/test/scala/com/twitter/logging/FileHandlerTest.scala b/util-logging/src/test/scala/com/twitter/logging/FileHandlerTest.scala index a617744658..f8a7f587b4 100644 --- a/util-logging/src/test/scala/com/twitter/logging/FileHandlerTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/FileHandlerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -17,334 +17,311 @@ package com.twitter.logging import java.io._ -import java.util.{Calendar, Date, logging => javalog} - -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import com.twitter.conversions.storage._ -import com.twitter.conversions.time._ -import com.twitter.util.{TempFolder, Time} -@RunWith(classOf[JUnitRunner]) -class FileHandlerTest extends WordSpec with TempFolder { - def reader(filename: String) = { +import java.util.Calendar +import java.util.Date +import java.util.UUID +import java.util.{logging => javalog} +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.conversions.DurationOps._ +import com.twitter.io.TempFolder +import com.twitter.util.Time +import java.nio.file.attribute.BasicFileAttributes +import java.nio.file.Files +import java.nio.file.Path +import java.nio.file.Paths +import java.util.function.BiPredicate +import org.scalatest.funsuite.AnyFunSuite + +class FileHandlerTest extends AnyFunSuite with TempFolder { + def reader(filename: String): BufferedReader = { new BufferedReader(new InputStreamReader(new FileInputStream(new File(folderName, filename)))) } - def writer(filename: String) = { + def writer(filename: String): OutputStreamWriter = { new OutputStreamWriter(new FileOutputStream(new File(folderName, filename)), "UTF-8") } - "FileHandler" should { - val record1 = new javalog.LogRecord(Level.INFO, "first post!") - val record2 = new javalog.LogRecord(Level.INFO, "second post") + private[this] val matcher = (name: String) => + new BiPredicate[Path, BasicFileAttributes] { + override def test(path: Path, attributes: BasicFileAttributes): Boolean = { + path.toFile.getName.contains(name) + } + } - "honor append setting on logfiles" in { - withTempFolder { - val f = writer("test.log") - f.write("hello!\n") - f.close + // Looks for files that contain `name` in the given `dirToSearch`. + // This method is used in place of setting the `user.dir` + // system property to ensure JDK11 compatibility. See + // https://bugs.java.com/bugdatabase/view_bug.do?bug_id=JDK-8202127 + private def filesContainingName(dirToSearch: String, name: String): Array[String] = + Files + .find( + Paths.get(dirToSearch), + 1, + matcher(name) + ).toArray.map(_.toString) + + val record1 = new javalog.LogRecord(Level.INFO, "first post!") + val record2 = new javalog.LogRecord(Level.INFO, "second post") + + test("FileHandler should honor append setting on logfiles") { + withTempFolder { + val f = writer("test.log") + f.write("hello!\n") + f.close() - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Hourly, - append = true, - formatter = BareFormatter - ).apply() + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.Hourly, + append = true, + formatter = BareFormatter + ).apply() - handler.publish(record1) + handler.publish(record1) - val f2 = reader("test.log") - assert(f2.readLine == "hello!") - } + val f2 = reader("test.log") + assert(f2.readLine == "hello!") + } + + withTempFolder { + val f = writer("test.log") + f.write("hello!\n") + f.close() + + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.Hourly, + append = false, + formatter = BareFormatter + ).apply() - withTempFolder { - val f = writer("test.log") - f.write("hello!\n") - f.close + handler.publish(record1) - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Hourly, - append = false, - formatter = BareFormatter - ).apply() + val f2 = reader("test.log") + assert(f2.readLine == "first post!") + } + } + + test("FileHandler should roll hourly logs on time") { + withTempFolder { + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.Hourly, + append = true, + formatter = BareFormatter + ).apply() + assert(handler.computeNextRollTime(1206769996722L).contains(1206770400000L)) + assert(handler.computeNextRollTime(1206770400000L).contains(1206774000000L)) + assert(handler.computeNextRollTime(1206774000001L).contains(1206777600000L)) + } + } + + test("FileHandler should roll weekly logs on time") { + withTempFolder { + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.Weekly(Calendar.SUNDAY), + append = true, + formatter = new Formatter(timezone = Some("GMT-7:00")) + ).apply() + assert(handler.computeNextRollTime(1250354734000L).contains(1250406000000L)) + assert(handler.computeNextRollTime(1250404734000L).contains(1250406000000L)) + assert(handler.computeNextRollTime(1250406001000L).contains(1251010800000L)) + assert(handler.computeNextRollTime(1250486000000L).contains(1251010800000L)) + assert(handler.computeNextRollTime(1250496000000L).contains(1251010800000L)) + } + } + // verify that at the proper time, the log file rolls and resets. + test("FileHandler should roll logs into new files") { + withTempFolder { + val handler = + new FileHandler(folderName + "/test.log", Policy.Hourly, true, -1, BareFormatter, None) + Time.withCurrentTimeFrozen { time => handler.publish(record1) + val date = new Date(Time.now.inMilliseconds) + + time.advance(1.hour) + + handler.publish(record2) + handler.close() - val f2 = reader("test.log") - assert(f2.readLine == "first post!") + assert(reader("test-" + handler.timeSuffix(date) + ".log").readLine == "first post!") + assert(reader("test.log").readLine == "second post") } } + } - // /* Test is commented out according to REPLA-618 */ - // - // "respond to a sighup to reopen a logfile with sun.misc" in { - // try { - // val signalClass = Class.forName("sun.misc.Signal") - // val sighup = signalClass.getConstructor(classOf[String]).newInstance("HUP").asInstanceOf[Object] - // val raiseMethod = signalClass.getMethod("raise", signalClass) - - // withTempFolder { - // val handler = FileHandler( - // filename = folderName + "/new.log", - // rollPolicy = Policy.SigHup, - // append = true, - // formatter = BareFormatter - // ).apply() - - // val logFile = new File(folderName, "new.log") - // logFile.renameTo(new File(folderName, "old.log")) - // handler.publish(record1) - - // raiseMethod.invoke(null, sighup) - - // val newLogFile = new File(folderName, "new.log") - // newLogFile.exists() should eventually(be_==(true)) - - // handler.publish(record2) - - // val oldReader = reader("old.log") - // assert(oldReader.readLine == "first post!") - // val newReader = reader("new.log") - // assert(newReader.readLine == "second post") - // } - // } catch { - // case ex: ClassNotFoundException => - // } - // } - - "roll logs on time" should { - "hourly" in { - withTempFolder { - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Hourly, - append = true, - formatter = BareFormatter - ).apply() - assert(handler.computeNextRollTime(1206769996722L) == Some(1206770400000L)) - assert(handler.computeNextRollTime(1206770400000L) == Some(1206774000000L)) - assert(handler.computeNextRollTime(1206774000001L) == Some(1206777600000L)) - } - } + test("FileHandler should keep no more than N log files around") { + withTempFolder { + assert(new File(folderName).list().length == 0) + + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.Hourly, + append = true, + rotateCount = 2, + formatter = BareFormatter + ).apply() + + handler.publish(record1) + assert(new File(folderName).list().length == 1) + handler.roll() + + handler.publish(record1) + assert(new File(folderName).list().length == 2) + handler.roll() + + handler.publish(record1) + assert(new File(folderName).list().length == 2) + handler.close() + } + } - "weekly" in { - withTempFolder { - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Weekly(Calendar.SUNDAY), - append = true, - formatter = new Formatter(timezone = Some("GMT-7:00")) - ).apply() - assert(handler.computeNextRollTime(1250354734000L) == Some(1250406000000L)) - assert(handler.computeNextRollTime(1250404734000L) == Some(1250406000000L)) - assert(handler.computeNextRollTime(1250406001000L) == Some(1251010800000L)) - assert(handler.computeNextRollTime(1250486000000L) == Some(1251010800000L)) - assert(handler.computeNextRollTime(1250496000000L) == Some(1251010800000L)) - } + test("FileHandler ignores the target filename despite shorter filenames") { + // even if the sort order puts the target filename before `rotateCount` other + // files, it should not be removed + withTempFolder { + assert(new File(folderName).list().length == 0) + val namePrefix = "test" + val name = namePrefix + ".log" + + val handler = FileHandler( + filename = folderName + "/" + name, + rollPolicy = Policy.Hourly, + append = true, + rotateCount = 1, + formatter = BareFormatter + ).apply() + + // create a file without the '.log' suffix, which will sort before the target + new File(folderName, namePrefix).createNewFile() + + def flush() = { + handler.publish(record1) + handler.roll() + assert(new File(folderName).list().length == 3) } + + // the target, 1 rotated file, and the short file should all remain + (1 to 5).foreach { _ => flush() } + val fileSet = new File(folderName).list().toSet + assert(fileSet.contains(name)) + assert(fileSet.contains(namePrefix)) } + } - // verify that at the proper time, the log file rolls and resets. - "roll logs into new files" in { - withTempFolder { - val handler = - new FileHandler(folderName + "/test.log", Policy.Hourly, true, -1, BareFormatter, None) - Time.withCurrentTimeFrozen { time => - handler.publish(record1) - val date = new Date(Time.now.inMilliseconds) + test("FileHandler correctly handles relative paths") { - time.advance(1.hour) + val userDir = System.getProperty("user.dir") + val filenamePrefix = UUID.randomUUID().toString - handler.publish(record2) - handler.close() + try { + val handler = FileHandler( + filename = filenamePrefix + ".log", // Note relative path! + rollPolicy = Policy.Hourly, + append = true, + rotateCount = 2, + formatter = BareFormatter + ).apply() - assert(reader("test-" + handler.timeSuffix(date) + ".log").readLine == "first post!") - assert(reader("test.log").readLine == "second post") - } - } + handler.publish(record1) + assert(filesContainingName(userDir, filenamePrefix).length == 1) + handler.roll() + + handler.publish(record1) + assert(filesContainingName(userDir, filenamePrefix).length == 2) + handler.roll() + + handler.publish(record1) + assert(filesContainingName(userDir, filenamePrefix).length == 2) + handler.close() + + } finally { + filesContainingName(userDir, filenamePrefix).foreach(f => new File(f).delete()) } + } - "keep no more than N log files around" in { - withTempFolder { - assert(new File(folderName).list().length == 0) + test("FileHandler should roll log files based on max size") { + withTempFolder { + // roll the log on the 3rd write. + val maxSize = record1.getMessage.length * 3 - 1 + + assert(new File(folderName).list().length == 0) - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Hourly, - append = true, - rotateCount = 2, - formatter = BareFormatter - ).apply() + val handler = FileHandler( + filename = folderName + "/test.log", + rollPolicy = Policy.MaxSize(maxSize.bytes), + append = true, + formatter = BareFormatter + ).apply() + // move time forward so the rotated logfiles will have distinct names. + Time.withCurrentTimeFrozen { time => + time.advance(1.second) + handler.publish(record1) + assert(new File(folderName).list().length == 1) + time.advance(1.second) handler.publish(record1) assert(new File(folderName).list().length == 1) - handler.roll() + time.advance(1.second) handler.publish(record1) assert(new File(folderName).list().length == 2) - handler.roll() - + time.advance(1.second) handler.publish(record1) assert(new File(folderName).list().length == 2) - handler.close() - } - } - - "ignores the target filename despite shorter filenames" in { - // even if the sort order puts the target filename before `rotateCount` other - // files, it should not be removed - withTempFolder { - assert(new File(folderName).list().length == 0) - val namePrefix = "test" - val name = namePrefix + ".log" - - val handler = FileHandler( - filename = folderName + "/" + name, - rollPolicy = Policy.Hourly, - append = true, - rotateCount = 1, - formatter = BareFormatter - ).apply() - - // create a file without the '.log' suffix, which will sort before the target - new File(folderName, namePrefix).createNewFile() - - def flush() = { - handler.publish(record1) - handler.roll() - assert(new File(folderName).list().length == 3) - } - // the target, 1 rotated file, and the short file should all remain - (1 to 5).foreach { _ => - flush() - } - val fileSet = new File(folderName).list().toSet - assert(fileSet.contains(name) == true) - assert(fileSet.contains(namePrefix) == true) + time.advance(1.second) + handler.publish(record1) + assert(new File(folderName).list().length == 3) } - } - "correctly handles relative paths" in { - withTempFolder { - // user.dir will be replaced with the temp folder, - // and will be restored when the test is complete - val wdir = System.getProperty("user.dir") - - try { - System.setProperty("user.dir", folderName) - - val handler = FileHandler( - filename = "test.log", // Note relative path! - rollPolicy = Policy.Hourly, - append = true, - rotateCount = 2, - formatter = BareFormatter - ).apply() - - handler.publish(record1) - assert(new File(folderName).list().length == 1) - handler.roll() - - handler.publish(record1) - assert(new File(folderName).list().length == 2) - handler.roll() - - handler.publish(record1) - assert(new File(folderName).list().length == 2) - handler.close() - } finally { - // restore user.dir to its original configuration - System.setProperty("user.dir", wdir) - } - - } + handler.close() } + } - "roll log files based on max size" in { - withTempFolder { - // roll the log on the 3rd write. - val maxSize = record1.getMessage.length * 3 - 1 - - assert(new File(folderName).list().length == 0) - - val handler = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.MaxSize(maxSize.bytes), - append = true, - formatter = BareFormatter - ).apply() - - // move time forward so the rotated logfiles will have distinct names. - Time.withCurrentTimeFrozen { time => - time.advance(1.second) - handler.publish(record1) - assert(new File(folderName).list().length == 1) - time.advance(1.second) - handler.publish(record1) - assert(new File(folderName).list().length == 1) - - time.advance(1.second) - handler.publish(record1) - assert(new File(folderName).list().length == 2) - time.advance(1.second) - handler.publish(record1) - assert(new File(folderName).list().length == 2) - - time.advance(1.second) - handler.publish(record1) - assert(new File(folderName).list().length == 3) - } - - handler.close() + /** + * This test mimics the scenario that, if log file exists, FileHandler should pick up it's size + * and roll over at the set limit, in order to prevent the file size from growing. + * The test succeeds if the roll over takes place as desired else fails. + */ + withTempFolder { + test("FileHandler should rollover at set limit when test is restarted, execution") { + val logLevel = Level.INFO + val fileSizeInMegaBytes: Long = 1 + val record1 = new javalog.LogRecord(logLevel, "Sending bytes to fill up file") + val filename = folderName + "/LogFileDir/testFileSize.log" + val rollPolicy = Policy.MaxSize(fileSizeInMegaBytes.megabytes) + val rotateCount = 8 + val append = true + val formatter = BareFormatter + + val handler = FileHandler(filename, rollPolicy, append, rotateCount, formatter).apply() + for (a <- 1 to 20000) { + handler.publish(record1) } - } + handler.close() - /** - * This test mimics the scenario that, if log file exists, FileHandler should pick up it's size - * and roll over at the set limit, in order to prevent the file size from growing. - * The test succeeds if the roll over takes place as desired else fails. - */ - withTempFolder { - "rollover at set limit when test is restarted, execution" in { - val logLevel = Level.INFO - val fileSizeInMegaBytes: Long = 1 - val record1 = new javalog.LogRecord(logLevel, "Sending bytes to fill up file") - val filename = folderName + "/LogFileDir/testFileSize.log" - val rollPolicy = Policy.MaxSize(fileSizeInMegaBytes.megabytes) - val rotateCount = 8 - val append = true - val formatter = BareFormatter - - val handler = FileHandler(filename, rollPolicy, append, rotateCount, formatter).apply() - for (a <- 1 to 20000) { - handler.publish(record1) - } - handler.close() - - val handler2 = FileHandler(filename, rollPolicy, append, rotateCount, formatter).apply() - for (a <- 1 to 20000) { - handler2.publish(record1) - } - handler2.close() - - def listLogFiles(dir: String): List[File] = { - val d = new File(dir) - if (d.exists && d.isDirectory) { - d.listFiles.filter(_.isFile).toList - } else { - List[File]() - } + val handler2 = FileHandler(filename, rollPolicy, append, rotateCount, formatter).apply() + for (a <- 1 to 20000) { + handler2.publish(record1) + } + handler2.close() + + def listLogFiles(dir: String): List[File] = { + val d = new File(dir) + if (d.exists && d.isDirectory) { + d.listFiles.filter(_.isFile).toList + } else { + List[File]() } + } - val files = listLogFiles(folderName + "/LogFileDir") - files.foreach { f: File => - val len = f.length().bytes - if (len > fileSizeInMegaBytes.megabytes) { - fail("Failed to roll over the log file") - } + val files = listLogFiles(folderName + "/LogFileDir") + files.foreach { (f: File) => + val len = f.length().bytes + if (len > fileSizeInMegaBytes.megabytes) { + fail("Failed to roll over the log file") } } } diff --git a/util-logging/src/test/scala/com/twitter/logging/FormatterTest.scala b/util-logging/src/test/scala/com/twitter/logging/FormatterTest.scala index 4798346276..c82e126ca0 100644 --- a/util-logging/src/test/scala/com/twitter/logging/FormatterTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/FormatterTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -17,15 +17,10 @@ package com.twitter.logging import java.util.{logging => javalog} +import com.twitter.conversions.StringOps._ +import org.scalatest.funsuite.AnyFunSuite -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -import com.twitter.conversions.string._ - -@RunWith(classOf[JUnitRunner]) -class FormatterTest extends WordSpec { +class FormatterTest extends AnyFunSuite { val basicFormatter = new Formatter val utcFormatter = new Formatter( @@ -70,149 +65,179 @@ class FormatterTest extends WordSpec { record4.setMillis(1206769996722L) record4.setThrown(noStackTraceException) - "Formatter" should { - "create a prefix" in { - assert( - basicFormatter.formatPrefix(Level.ERROR, "20080329-05:53:16.722", "(root)") == - "ERR [20080329-05:53:16.722] (root): " - ) - assert( - basicFormatter.formatPrefix(Level.DEBUG, "20080329-05:53:16.722", "(root)") == - "DEB [20080329-05:53:16.722] (root): " - ) - assert( - basicFormatter.formatPrefix(Level.WARNING, "20080329-05:53:16.722", "(root)") == - "WAR [20080329-05:53:16.722] (root): " - ) - } + test("Formatter should create a prefix") { + assert( + basicFormatter.formatPrefix(Level.ERROR, "20080329-05:53:16.722", "(root)") == + "ERR [20080329-05:53:16.722] (root): " + ) + assert( + basicFormatter.formatPrefix(Level.DEBUG, "20080329-05:53:16.722", "(root)") == + "DEB [20080329-05:53:16.722] (root): " + ) + assert( + basicFormatter.formatPrefix(Level.WARNING, "20080329-05:53:16.722", "(root)") == + "WAR [20080329-05:53:16.722] (root): " + ) + } - "format a log level name" in { - assert(basicFormatter.formatLevelName(Level.ERROR) == "ERROR") - assert(basicFormatter.formatLevelName(Level.DEBUG) == "DEBUG") - assert(basicFormatter.formatLevelName(Level.WARNING) == "WARNING") - } + test("Formatter should format a log level name") { + assert(basicFormatter.formatLevelName(Level.ERROR) == "ERROR") + assert(basicFormatter.formatLevelName(Level.DEBUG) == "DEBUG") + assert(basicFormatter.formatLevelName(Level.WARNING) == "WARNING") + } - "format text" in { - val record = new LogRecord(Level.ERROR, "error %s") - assert(basicFormatter.formatText(record) == "error %s") - record.setParameters(Array("123")) - assert(basicFormatter.formatText(record) == "error 123") - } + test("Formatter should format text") { + val record = new LogRecord(Level.ERROR, "error %s") + assert(basicFormatter.formatText(record) == "error %s") + record.setParameters(Array("123")) + assert(basicFormatter.formatText(record) == "error 123") + } - "format a timestamp" in { - assert(utcFormatter.format(record1) == "ERR [20080329-05:53:16.722] jobs: boo.\n") - } + test("Formatter should format message lines") { + val record = new LogRecord(Level.ERROR, "how now brown cow") + assert(basicFormatter.formatMessageLines(record) sameElements Array("how now brown cow")) + } - "do lazy message evaluation" in { - var callCount = 0 - def getSideEffect = { - callCount += 1 - "ok" - } + test("Formatter should format message lines multiple values") { + val record = new LogRecord(Level.ERROR, "how now brown cow\nthe quick brown fox") + assert( + basicFormatter + .formatMessageLines(record) sameElements Array("how now brown cow", "the quick brown fox")) + } - val record = new LazyLogRecord(Level.DEBUG, "this is " + getSideEffect) + test("Formatter should format message lines with message") { + val record = new LogRecord(Level.ERROR, "how now brown cow") + assert( + basicFormatter.formatMessageLines("how now brown cow", record.getThrown) sameElements Array( + "how now brown cow")) + } - assert(callCount == 0) - assert(basicFormatter.formatText(record) == "this is ok") - assert(callCount == 1) - assert(basicFormatter.formatText(record) == "this is ok") - assert(callCount == 1) - } + test("Formatter should format message lines with message and multiple values") { + val record = new LogRecord(Level.ERROR, "how now brown cow\nthe quick brown fox") + // ignores the message from the record + assert( + basicFormatter.formatMessageLines("how now brown cow", record.getThrown) sameElements Array( + "how now brown cow")) + } - "format package names" in { - assert(utcFormatter.format(record1) == "ERR [20080329-05:53:16.722] jobs: boo.\n") - assert( - fullPackageFormatter.format(record1) == - "ERR [20080329-05:53:16.722] com.example.jobs: boo.\n" - ) - } + test("Formatter should format message lines with message and multiple values again") { + val record = new LogRecord(Level.ERROR, "how now brown cow\nthe quick brown fox") + // ignores the message from the record + assert( + basicFormatter + .formatMessageLines( + "how now brown cow\nthe quick brown fox", + record.getThrown) sameElements Array("how now brown cow", "the quick brown fox")) + } - "handle other prefixes" in { - assert(prefixFormatter.format(record2) == "jobs 05:53 DEBU useless info.\n") - } + test("Formatter should format a timestamp") { + assert(utcFormatter.format(record1) == "ERR [20080329-05:53:16.722] jobs: boo.\n") + } - "truncate line" in { - assert( - truncateFormatter.format(record3) == - "CRI [20080329-05:53:16.722] whiskey: Something terrible happened th...\n" - ) + test("Formatter should do lazy message evaluation") { + var callCount = 0 + def getSideEffect = { + callCount += 1 + "ok" } - "write stack traces" should { - object ExceptionLooper { - def cycle(n: Int): Unit = { - if (n == 0) { - throw new Exception("Aie!") - } else { - cycle(n - 1) - throw new Exception("this is just here to fool the tail recursion optimizer") - } - } - - def cycle2(n: Int): Unit = { - try { - cycle(n) - } catch { - case t: Throwable => - throw new Exception("grrrr", t) - } - } - } + val record = new LazyLogRecord(Level.DEBUG, "this is " + getSideEffect) + + assert(callCount == 0) + assert(basicFormatter.formatText(record) == "this is ok") + assert(callCount == 1) + assert(basicFormatter.formatText(record) == "this is ok") + assert(callCount == 1) + } + + test("Formatter should format package names") { + assert(utcFormatter.format(record1) == "ERR [20080329-05:53:16.722] jobs: boo.\n") + assert( + fullPackageFormatter.format(record1) == + "ERR [20080329-05:53:16.722] com.example.jobs: boo.\n" + ) + } + + test("Formatter should handle other prefixes") { + assert(prefixFormatter.format(record2) == "jobs 05:53 DEBU useless info.\n") + } + + test("Formatter should truncate line") { + assert( + truncateFormatter.format(record3) == + "CRI [20080329-05:53:16.722] whiskey: Something terrible happened th...\n" + ) + } - def scrub(in: String) = { - in.regexSub("""FormatterTest.scala:\d+""".r) { m => - "FormatterTest.scala:NNN" - } - .regexSub("""FormatterTest\$[\w\\$]+""".r) { m => - "FormatterTest$$" - } + object ExceptionLooper { + def cycle(n: Int): Unit = { + if (n == 0) { + throw new Exception("Aie!") + } else { + cycle(n - 1) + throw new Exception("this is just here to fool the tail recursion optimizer") } + } - "simple" in { - val exception = try { - ExceptionLooper.cycle(10) - null - } catch { - case t: Throwable => t - } - assert( - Formatter.formatStackTrace(exception, 5).map { scrub(_) } == List( - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " (...more...)" - ) - ) + def cycle2(n: Int): Unit = { + try { + cycle(n) + } catch { + case t: Throwable => + throw new Exception("grrrr", t) } + } + } - "nested" in { - val exception = try { - ExceptionLooper.cycle2(2) - null - } catch { - case t: Throwable => t - } - assert( - Formatter.formatStackTrace(exception, 1).map { scrub(_) } == List( - " at com.twitter.logging.FormatterTest$$.cycle2(FormatterTest.scala:NNN)", - " (...more...)", - "Caused by java.lang.Exception: Aie!", - " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", - " (...more...)" - ) - ) + def scrub(in: String) = { + in.regexSub("""FormatterTest.scala:\d+""".r) { m => "FormatterTest.scala:NNN" } + .regexSub("""FormatterTest\$[\w\\$]+""".r) { m => "FormatterTest$$" } + } + test("Formatter should write simple stack traces") { + val exception = + try { + ExceptionLooper.cycle(10) + null + } catch { + case t: Throwable => t } + assert( + Formatter.formatStackTrace(exception, 5).map { scrub(_) } == List( + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " (...more...)" + ) + ) + } - "log even blank exceptions" in { - assert( - utcFormatter.format(record4) == - "ERR [20080329-05:53:16.722] jobs: with minimal exception\n" + - "ERR [20080329-05:53:16.722] jobs: java.lang.Exception: fast exception no stacktrace\n" - ) + test("Formatter should write nested stack traces") { + val exception = + try { + ExceptionLooper.cycle2(2) + null + } catch { + case t: Throwable => t } - } + assert( + Formatter.formatStackTrace(exception, 1).map { scrub(_) } == List( + " at com.twitter.logging.FormatterTest$$.cycle2(FormatterTest.scala:NNN)", + " (...more...)", + "Caused by java.lang.Exception: Aie!", + " at com.twitter.logging.FormatterTest$$.cycle(FormatterTest.scala:NNN)", + " (...more...)" + ) + ) + } + + test("Formatter should log even blank exceptions") { + assert( + utcFormatter.format(record4) == + "ERR [20080329-05:53:16.722] jobs: with minimal exception\n" + + "ERR [20080329-05:53:16.722] jobs: java.lang.Exception: fast exception no stacktrace\n" + ) } } diff --git a/util-logging/src/test/scala/com/twitter/logging/HasLogLevelTest.scala b/util-logging/src/test/scala/com/twitter/logging/HasLogLevelTest.scala index df570d8270..f7eecd2534 100644 --- a/util-logging/src/test/scala/com/twitter/logging/HasLogLevelTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/HasLogLevelTest.scala @@ -1,11 +1,8 @@ package com.twitter.logging -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class HasLogLevelTest extends FunSuite { +class HasLogLevelTest extends AnyFunSuite { private class WithLogLevel(val logLevel: Level, cause: Throwable = null) extends Exception(cause) diff --git a/util-logging/src/test/scala/com/twitter/logging/LogRecordTest.scala b/util-logging/src/test/scala/com/twitter/logging/LogRecordTest.scala index e83febd379..a9ac801d9f 100644 --- a/util-logging/src/test/scala/com/twitter/logging/LogRecordTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/LogRecordTest.scala @@ -1,17 +1,12 @@ package com.twitter.logging import java.util.logging.{LogRecord => JRecord} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class LogRecordTest extends FunSuite { +class LogRecordTest extends AnyFunSuite { test("LogRecord should getMethod properly") { Logger.withLoggers(Nil) { - new LogRecordTestHelper({ r: JRecord => - r.getSourceMethodName() - }) { + new LogRecordTestHelper({ (r: JRecord) => r.getSourceMethodName() }) { def makingLogRecord() = { logger.log(Level.INFO, "OK") assert(handler.get == "makingLogRecord") @@ -39,10 +34,7 @@ abstract class LogRecordTestHelper(formats: JRecord => String) { logger.addHandler(handler) } -class Foo - extends LogRecordTestHelper({ r: JRecord => - r.getSourceClassName() - }) { +class Foo extends LogRecordTestHelper({ (r: JRecord) => r.getSourceClassName }) { def makingLogRecord(): Unit = { logger.log(Level.INFO, "OK") } diff --git a/util-logging/src/test/scala/com/twitter/logging/LoggerTest.scala b/util-logging/src/test/scala/com/twitter/logging/LoggerTest.scala index 0c664bab42..ab849d7c3e 100644 --- a/util-logging/src/test/scala/com/twitter/logging/LoggerTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/LoggerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,18 +16,20 @@ package com.twitter.logging -import com.twitter.conversions.time._ -import com.twitter.util.TempFolder +import com.twitter.conversions.DurationOps._ +import com.twitter.io.TempFolder import java.net.InetSocketAddress -import java.util.concurrent.{Callable, CountDownLatch, Executors, Future, TimeUnit} +import java.util.concurrent.Callable +import java.util.concurrent.CountDownLatch +import java.util.concurrent.Executors +import java.util.concurrent.Future +import java.util.concurrent.TimeUnit import java.util.{logging => javalog} -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.{BeforeAndAfter, WordSpec} +import org.scalatest.BeforeAndAfter +import org.scalatest.funsuite.AnyFunSuite import scala.collection.mutable -@RunWith(classOf[JUnitRunner]) -class LoggerTest extends WordSpec with TempFolder with BeforeAndAfter { +class LoggerTest extends AnyFunSuite with TempFolder with BeforeAndAfter { val logLevel = Logger.levelNames(Option[String](System.getenv("log")).getOrElse("FATAL").toUpperCase) @@ -70,7 +72,7 @@ class LoggerTest extends WordSpec with TempFolder with BeforeAndAfter { logger.addHandler(traceHandler) } - def logLines(): Seq[String] = traceHandler.get.split("\n") + def logLines(): Seq[String] = traceHandler.get.split("\n").toIndexedSeq /** * Verify that the logger set up with `traceLogger` has received a log line with the given @@ -105,369 +107,349 @@ class LoggerTest extends WordSpec with TempFolder with BeforeAndAfter { case class WithLogLevel(logLevel: Level) extends Exception("") with HasLogLevel - "Logger" should { - val h = new LoggerSpecHelper - import h._ + val h = new LoggerSpecHelper + import h._ - def before() = { - Logger.clearHandlers() - timeFrozenHandler.clear() - myHandler = new StringHandler(BareFormatter, None) - log = Logger.get("") - log.setLevel(Level.ERROR) - log.addHandler(myHandler) - } - - "not execute string creation and concatenation" in { - val logger = Logger.get("lazyTest1") - var executed = false - def function() = { - executed = true - "asdf" + executed + " hi there" - } + def before() = { + Logger.clearHandlers() + timeFrozenHandler.clear() + myHandler = new StringHandler(BareFormatter, None) + log = Logger.get("") + log.setLevel(Level.ERROR) + log.addHandler(myHandler) + } - logger.debugLazy(function) - assert(!executed) + test("Logger should not execute string creation and concatenation") { + val logger = Logger.get("lazyTest1") + var executed = false + def function() = { + executed = true + "asdf" + executed + " hi there" } - "execute string creation and concatenation is FAILING" in { - val logger = Logger.get("lazyTest2") - logger.setLevel(Level.DEBUG) - var executed = false - def function() = { - executed = true - "asdf" + executed + " hi there" - } - logger.debugLazy(function) - assert(executed) - } + logger.debugLazy(function()) + assert(!executed) + } - "make sure compiles with normal string case" in { - //in some cases, if debugLazy definition was changed, the below would no longer compile - val logger = Logger.get("lazyTest3") - val executed = true - logger.debugLazy("hi there" + executed + "cool") + test("Logger should execute string creation and concatenation is FAILING") { + val logger = Logger.get("lazyTest2") + logger.setLevel(Level.DEBUG) + var executed = false + def function() = { + executed = true + "asdf" + executed + " hi there" } + logger.debugLazy(function()) + assert(executed) + } - "provide level name and value maps" in { - assert( - Logger.levels == Map( - Level.ALL.value -> Level.ALL, - Level.TRACE.value -> Level.TRACE, - Level.DEBUG.value -> Level.DEBUG, - Level.INFO.value -> Level.INFO, - Level.WARNING.value -> Level.WARNING, - Level.ERROR.value -> Level.ERROR, - Level.CRITICAL.value -> Level.CRITICAL, - Level.FATAL.value -> Level.FATAL, - Level.OFF.value -> Level.OFF - ) + test("Logger should make sure compiles with normal string case") { + //in some cases, if debugLazy definition was changed, the below would no longer compile + val logger = Logger.get("lazyTest3") + val executed = true + logger.debugLazy("hi there" + executed + "cool") + } + + test("Logger should provide level name and value maps") { + assert( + Logger.levels == Map( + Level.ALL.value -> Level.ALL, + Level.TRACE.value -> Level.TRACE, + Level.DEBUG.value -> Level.DEBUG, + Level.INFO.value -> Level.INFO, + Level.WARNING.value -> Level.WARNING, + Level.ERROR.value -> Level.ERROR, + Level.CRITICAL.value -> Level.CRITICAL, + Level.FATAL.value -> Level.FATAL, + Level.OFF.value -> Level.OFF ) - assert( - Logger.levelNames == Map( - "ALL" -> Level.ALL, - "TRACE" -> Level.TRACE, - "DEBUG" -> Level.DEBUG, - "INFO" -> Level.INFO, - "WARNING" -> Level.WARNING, - "ERROR" -> Level.ERROR, - "CRITICAL" -> Level.CRITICAL, - "FATAL" -> Level.FATAL, - "OFF" -> Level.OFF - ) + ) + assert( + Logger.levelNames == Map( + "ALL" -> Level.ALL, + "TRACE" -> Level.TRACE, + "DEBUG" -> Level.DEBUG, + "INFO" -> Level.INFO, + "WARNING" -> Level.WARNING, + "ERROR" -> Level.ERROR, + "CRITICAL" -> Level.CRITICAL, + "FATAL" -> Level.FATAL, + "OFF" -> Level.OFF ) - } + ) + } - "figure out package names" in { - val log1 = Logger(this.getClass) - assert(log1.name == "com.twitter.logging.LoggerTest") - } + test("Logger should figure out package names") { + val log1 = Logger(this.getClass) + assert(log1.name == "com.twitter.logging.LoggerTest") + } - "log & trace a message" in { - traceLogger(Level.INFO) - Logger.get("").info("angry duck") - mustLog("duck") - } + test("Logger should log & trace a message") { + traceLogger(Level.INFO) + Logger.get("").info("angry duck") + mustLog("duck") + } - "log & trace a throwable with a log level" in { - traceLogger(Level.WARNING) - Logger.get("").throwable(WithLogLevel(Level.WARNING), "bananas") - Logger.get("").throwable(WithLogLevel(Level.INFO), "apples") - mustLog("bananas") - mustNotLog("apples") - } + test("Logger should log & trace a throwable with a log level") { + traceLogger(Level.WARNING) + Logger.get("").throwable(WithLogLevel(Level.WARNING), "bananas") + Logger.get("").throwable(WithLogLevel(Level.INFO), "apples") + mustLog("bananas") + mustNotLog("apples") + } - "log & trace a throwable with a default log level" in { - traceLogger(Level.ERROR) - Logger.get("").throwable(new Exception(), "angry duck") - Logger.get("").throwable(new Exception(), "invisible", defaultLevel = Level.INFO) - mustLog("duck") - mustNotLog("invisible") - } + test("Logger should log & trace a throwable with a default log level") { + traceLogger(Level.ERROR) + Logger.get("").throwable(new Exception(), "angry duck") + Logger.get("").throwable(new Exception(), "invisible", defaultLevel = Level.INFO) + mustLog("duck") + mustNotLog("invisible") + } - "log & trace a message with percent signs" in { - traceLogger(Level.INFO) - Logger.get("")(Level.INFO, "%i") - mustLog("%i") - } + test("Logger should log & trace a message with percent signs") { + traceLogger(Level.INFO) + Logger.get("")(Level.INFO, "%i") + Logger.get("")(Level.INFO, new Exception(), "%j") + mustLog("%i") + mustLog("%j") + } - "log a message, with timestamp" in { - before() + test("Logger should log a message, with timestamp") { + before() - Logger.clearHandlers() - myHandler = timeFrozenHandler - log.addHandler(timeFrozenHandler) - log.error("error!") - assert(parse() == List("ERR [20080329-05:53:16.722] (root): error!")) - } + Logger.clearHandlers() + myHandler = timeFrozenHandler + log.addHandler(timeFrozenHandler) + log.error("error!") + assert(parse() == List("ERR [20080329-05:53:16.722] (root): error!")) + } - "get single-threaded return the same value" in { - val loggerFirst = Logger.get("getTest") - assert(loggerFirst != null) + test("Logger should get single-threaded return the same value") { + val loggerFirst = Logger.get("getTest") + assert(loggerFirst != null) - val loggerSecond = Logger.get("getTest") - assert(loggerSecond == loggerFirst) + val loggerSecond = Logger.get("getTest") + assert(loggerSecond == loggerFirst) + } + + test("Logger should get multi-threaded return the same value") { + val numThreads = 10 + val latch = new CountDownLatch(1) + + // queue up the workers + val executorService = Executors.newFixedThreadPool(numThreads) + val futureResults = new mutable.ListBuffer[Future[Logger]] + for (i <- 0.until(numThreads)) { + val future = executorService.submit(new Callable[Logger]() { + def call(): Logger = { + latch.await(10, TimeUnit.SECONDS) + return Logger.get("concurrencyTest") + } + }) + futureResults += future } + executorService.shutdown + // let them rip, and then wait for em to finish + latch.countDown + assert(executorService.awaitTermination(10, TimeUnit.SECONDS) == true) + + // now make sure they are all the same reference + val expected = futureResults(0).get + for (i <- 1.until(numThreads)) { + val result = futureResults(i).get + assert(result == expected) + } + } - "get multi-threaded return the same value" in { - val numThreads = 10 - val latch = new CountDownLatch(1) - - // queue up the workers - val executorService = Executors.newFixedThreadPool(numThreads) - val futureResults = new mutable.ListBuffer[Future[Logger]] - for (i <- 0.until(numThreads)) { - val future = executorService.submit(new Callable[Logger]() { - def call(): Logger = { - latch.await(10, TimeUnit.SECONDS) - return Logger.get("concurrencyTest") - } - }) - futureResults += future - } - executorService.shutdown - // let them rip, and then wait for em to finish - latch.countDown - assert(executorService.awaitTermination(10, TimeUnit.SECONDS) == true) - - // now make sure they are all the same reference - val expected = futureResults(0).get - for (i <- 1.until(numThreads)) { - val result = futureResults(i).get - assert(result == expected) - } + test( + "Logger should withLoggers applies logger factories, executes a block, and then applies original factories") { + val initialFactories = List(LoggerFactory(node = "", level = Some(Level.DEBUG))) + val otherFactories = List(LoggerFactory(node = "", level = Some(Level.INFO))) + Logger.configure(initialFactories) + + assert(Logger.get("").getLevel() == Level.DEBUG) + Logger.withLoggers(otherFactories) { + assert(Logger.get("").getLevel() == Level.INFO) } + assert(Logger.get("").getLevel() == Level.DEBUG) + } - "withLoggers applies logger factories, executes a block, and then applies original factories" in { - val initialFactories = List(LoggerFactory(node = "", level = Some(Level.DEBUG))) - val otherFactories = List(LoggerFactory(node = "", level = Some(Level.INFO))) - Logger.configure(initialFactories) + test("Logger configure logging should") { + def before(): Unit = { + Logger.clearHandlers() + } + } - assert(Logger.get("").getLevel == Level.DEBUG) - Logger.withLoggers(otherFactories) { - assert(Logger.get("").getLevel() == Level.INFO) - } - assert(Logger.get("").getLevel == Level.DEBUG) + test("Logger configure logging with file handler") { + val separator = java.io.File.pathSeparator + withTempFolder { + val log: Logger = LoggerFactory( + node = "com.twitter", + level = Some(Level.DEBUG), + handlers = FileHandler( + filename = folderName + separator + "test.log", + rollPolicy = Policy.Never, + append = false, + level = Some(Level.INFO), + formatter = new Formatter( + useFullPackageNames = true, + truncateAt = 1024, + prefix = "%s %s" + ) + ) :: Nil + ).apply() + + assert(log.getLevel() == Level.DEBUG) + assert(log.getHandlers().length == 1) + val handler = log.getHandlers()(0).asInstanceOf[FileHandler] + val fileName = folderName + separator + "test.log" + assert(handler.filename == fileName) + assert(handler.append == false) + assert(handler.getLevel() == Level.INFO) + val formatter = handler.formatter + assert( + formatter.formatPrefix(javalog.Level.WARNING, "10:55", "hello") == "WARNING 10:55 hello" + ) + assert(log.name == "com.twitter") + assert(formatter.truncateAt == 1024) + assert(formatter.useFullPackageNames == true) } + } - "configure logging" should { - def before(): Unit = { - Logger.clearHandlers() - } + test("Logger configure logging with syslog handler") { + withTempFolder { + val log: Logger = LoggerFactory( + node = "com.twitter", + handlers = SyslogHandler( + formatter = new SyslogFormatter( + serverName = Some("elmo"), + priority = 128 + ), + server = "localhost", + port = 212 + ) :: Nil + ).apply() + + assert(log.getHandlers().length == 1) + val h = log.getHandlers()(0).asInstanceOf[SyslogHandler] + assert(h.dest.asInstanceOf[InetSocketAddress].getHostName == "localhost") + assert(h.dest.asInstanceOf[InetSocketAddress].getPort == 212) + val formatter = h.formatter.asInstanceOf[SyslogFormatter] + assert(formatter.serverName == Some("elmo")) + assert(formatter.priority == 128) + } + } - "file handler" in { - withTempFolder { - val log: Logger = LoggerFactory( - node = "com.twitter", - level = Some(Level.DEBUG), - handlers = FileHandler( - filename = folderName + "/test.log", - rollPolicy = Policy.Never, - append = false, - level = Some(Level.INFO), - formatter = new Formatter( - useFullPackageNames = true, - truncateAt = 1024, - prefix = "%s %s" - ) - ) :: Nil - ).apply() - - assert(log.getLevel == Level.DEBUG) - assert(log.getHandlers().length == 1) - val handler = log.getHandlers()(0).asInstanceOf[FileHandler] - assert(handler.filename == folderName + "/test.log") - assert(handler.append == false) - assert(handler.getLevel == Level.INFO) - val formatter = handler.formatter - assert( - formatter.formatPrefix(javalog.Level.WARNING, "10:55", "hello") == "WARNING 10:55 hello" + test("Logger configure logging with complex config") { + withTempFolder { + val factories = LoggerFactory( + level = Some(Level.INFO), + handlers = ThrottledHandler( + duration = 60.seconds, + maxToDisplay = 10, + handler = FileHandler( + filename = folderName + "/production.log", + rollPolicy = Policy.SigHup, + formatter = new Formatter( + truncateStackTracesAt = 100 + ) ) - assert(log.name == "com.twitter") - assert(formatter.truncateAt == 1024) - assert(formatter.useFullPackageNames == true) - } + ) :: Nil + ) :: LoggerFactory( + node = "w3c", + level = Some(Level.OFF), + useParents = false + ) :: LoggerFactory( + node = "bad_jobs", + level = Some(Level.INFO), + useParents = false, + handlers = FileHandler( + filename = folderName + "/bad_jobs.log", + rollPolicy = Policy.Never + ) :: Nil + ) :: Nil + + Logger.configure(factories) + assert(Logger.get("").getLevel() == Level.INFO) + assert(Logger.get("w3c").getLevel() == Level.OFF) + assert(Logger.get("bad_jobs").getLevel() == Level.INFO) + try { + Logger.get("").getHandlers()(0).asInstanceOf[ThrottledHandler] + } catch { + case _: ClassCastException => fail("not a ThrottledHandler") } - - "syslog handler" in { - withTempFolder { - val log: Logger = LoggerFactory( - node = "com.twitter", - handlers = SyslogHandler( - formatter = new SyslogFormatter( - serverName = Some("elmo"), - priority = 128 - ), - server = "localhost", - port = 212 - ) :: Nil - ).apply() - - assert(log.getHandlers.length == 1) - val h = log.getHandlers()(0).asInstanceOf[SyslogHandler] - assert(h.dest.asInstanceOf[InetSocketAddress].getHostName == "localhost") - assert(h.dest.asInstanceOf[InetSocketAddress].getPort == 212) - val formatter = h.formatter.asInstanceOf[SyslogFormatter] - assert(formatter.serverName == Some("elmo")) - assert(formatter.priority == 128) - } + try { + Logger + .get("") + .getHandlers()(0) + .asInstanceOf[ThrottledHandler] + .handler + .asInstanceOf[FileHandler] + } catch { + case _: ClassCastException => fail("not a FileHandler") } - - "complex config" in { - withTempFolder { - val factories = LoggerFactory( - level = Some(Level.INFO), - handlers = ThrottledHandler( - duration = 60.seconds, - maxToDisplay = 10, - handler = FileHandler( - filename = folderName + "/production.log", - rollPolicy = Policy.SigHup, - formatter = new Formatter( - truncateStackTracesAt = 100 - ) - ) - ) :: Nil - ) :: LoggerFactory( - node = "w3c", - level = Some(Level.OFF), - useParents = false - ) :: LoggerFactory( - node = "stats", - level = Some(Level.INFO), - useParents = false, - handlers = ScribeHandler( - formatter = BareFormatter, - maxMessagesToBuffer = 100, - category = "cuckoo_json" - ) :: Nil - ) :: LoggerFactory( - node = "bad_jobs", - level = Some(Level.INFO), - useParents = false, - handlers = FileHandler( - filename = folderName + "/bad_jobs.log", - rollPolicy = Policy.Never - ) :: Nil - ) :: Nil - - Logger.configure(factories) - assert(Logger.get("").getLevel == Level.INFO) - assert(Logger.get("w3c").getLevel == Level.OFF) - assert(Logger.get("stats").getLevel == Level.INFO) - assert(Logger.get("bad_jobs").getLevel == Level.INFO) - try { - Logger.get("").getHandlers()(0).asInstanceOf[ThrottledHandler] - } catch { - case _: ClassCastException => fail("not a ThrottledHandler") - } - try { - Logger - .get("") - .getHandlers()(0) - .asInstanceOf[ThrottledHandler] - .handler - .asInstanceOf[FileHandler] - } catch { - case _: ClassCastException => fail("not a FileHandler") - } - assert(Logger.get("w3c").getHandlers().size == 0) - try { - Logger.get("stats").getHandlers()(0).asInstanceOf[ScribeHandler] - } catch { - case _: ClassCastException => fail("not a ScribeHandler") - } - try { - Logger.get("bad_jobs").getHandlers()(0).asInstanceOf[FileHandler] - } catch { - case _: ClassCastException => fail("not a FileHandler") - } - - } + assert(Logger.get("w3c").getHandlers().size == 0) + try { + Logger.get("bad_jobs").getHandlers()(0).asInstanceOf[FileHandler] + } catch { + case _: ClassCastException => fail("not a FileHandler") } - } - "java logging" should { - val logger = javalog.Logger.getLogger("") - - def before() = { - traceLogger(Level.INFO) - } + } + } - "single arg calls" in { - before() - logger.log(javalog.Level.INFO, "V1={0}", "A") - mustLog("V1=A") - } + test("java logging single arg calls") { + val logger = javalog.Logger.getLogger("") + traceLogger(Level.INFO) + logger.log(javalog.Level.INFO, "V1={0}", "A") + mustLog("V1=A") + } - "varargs calls" in { - before() - logger.log(javalog.Level.INFO, "V1={0}, V2={1}", Array[AnyRef]("A", "B")) - mustLog("V1=A, V2=B") - } + test("java logging varargs calls") { + val logger = javalog.Logger.getLogger("") + traceLogger(Level.INFO) + logger.log(javalog.Level.INFO, "V1={0}, V2={1}", Array[AnyRef]("A", "B")) + mustLog("V1=A, V2=B") + } - "invalid message format" in { - before() - logger.log(javalog.Level.INFO, "V1=%s", "A") - mustLog("V1=%s") // %s notation is not known in java MessageFormat - } + test("java logging invalid message format") { + val logger = javalog.Logger.getLogger("") + traceLogger(Level.INFO) + logger.log(javalog.Level.INFO, "V1=%s", "A") + mustLog("V1=%s") // %s notation is not known)java MessageFormat + } - // logging in scala uses the %s format and not the Java MessageFormat - "compare scala logging format" in { - before() - Logger.get("").info("V1{0}=%s", "A") - mustLog("V1{0}=A") - } - } + // logging in scala uses the %s format and not the Java MessageFormat + test("java logging compare scala logging format") { + val logger = javalog.Logger.getLogger("") + traceLogger(Level.INFO) + Logger.get("").info("V1{0}=%s", "A") + mustLog("V1{0}=A") } - "Levels.fromJava" should { - "return a corresponding level if it exists" in { - assert(Level.fromJava(javalog.Level.OFF) == Some(Level.OFF)) - assert(Level.fromJava(javalog.Level.SEVERE) == Some(Level.FATAL)) - // No corresponding level for ERROR and CRITICAL in javalog. - assert(Level.fromJava(javalog.Level.WARNING) == Some(Level.WARNING)) - assert(Level.fromJava(javalog.Level.INFO) == Some(Level.INFO)) - assert(Level.fromJava(javalog.Level.FINE) == Some(Level.DEBUG)) - assert(Level.fromJava(javalog.Level.FINER) == Some(Level.TRACE)) - assert(Level.fromJava(javalog.Level.ALL) == Some(Level.ALL)) - } + test("Levels.fromJava should return a corresponding level if it exists") { + assert(Level.fromJava(javalog.Level.OFF) == Some(Level.OFF)) + assert(Level.fromJava(javalog.Level.SEVERE) == Some(Level.FATAL)) + // No corresponding level for ERROR and CRITICAL in javalog. + assert(Level.fromJava(javalog.Level.WARNING) == Some(Level.WARNING)) + assert(Level.fromJava(javalog.Level.INFO) == Some(Level.INFO)) + assert(Level.fromJava(javalog.Level.FINE) == Some(Level.DEBUG)) + assert(Level.fromJava(javalog.Level.FINER) == Some(Level.TRACE)) + assert(Level.fromJava(javalog.Level.ALL) == Some(Level.ALL)) + } - "return None for non corresponding levels" in { - assert(Level.fromJava(javalog.Level.CONFIG) == None) - assert(Level.fromJava(javalog.Level.FINEST) == None) - } + test("Levels.fromJava should return None for non corresponding levels") { + assert(Level.fromJava(javalog.Level.CONFIG) == None) + assert(Level.fromJava(javalog.Level.FINEST) == None) } - "Levels.parse" should { - "return Levels for valid names" in { - assert(Level.parse("INFO") == Some(Level.INFO)) - assert(Level.parse("TRACE") == Some(Level.TRACE)) - } + test("Levels.parse should return Levels for valid names") { + assert(Level.parse("INFO") == Some(Level.INFO)) + assert(Level.parse("TRACE") == Some(Level.TRACE)) + } - "return None for invalid names" in { - assert(Level.parse("") == None) - assert(Level.parse("FOO") == None) - } + test("Levels.parse should return None for invalid names") { + assert(Level.parse("") == None) + assert(Level.parse("FOO") == None) } } diff --git a/util-logging/src/test/scala/com/twitter/logging/LoggingTest.scala b/util-logging/src/test/scala/com/twitter/logging/LoggingTest.scala deleted file mode 100644 index 394e034641..0000000000 --- a/util-logging/src/test/scala/com/twitter/logging/LoggingTest.scala +++ /dev/null @@ -1,59 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter.logging - -import com.twitter.app.App -import org.scalatest.{BeforeAndAfterEach, BeforeAndAfterAll, FunSuite} - -class LoggingTest extends FunSuite with BeforeAndAfterEach with BeforeAndAfterAll { - // ensure this is initialized--we need to do something with it or scala - // complains about purity - require(ScribeHandler.log != null) - - object TestLoggingApp extends App with Logging { - override def handlers = ScribeHandler() :: super.handlers - } - - object TestLoggingWithConfigureApp extends App with Logging { - override def handlers = ScribeHandler() :: super.handlers - override def configureLoggerFactories(): Unit = {} - } - - test("TestLoggingApp should have one factory with two log handlers") { - TestLoggingApp.main(Array.empty) - assert(TestLoggingApp.loggerFactories.size == 1) - assert(TestLoggingApp.loggerFactories.head.handlers.size == 2) - // Logger is configured with two handlers, and has one logger - assert(Logger.iterator.size == 1) - } - - test("TestLoggingWithConfigureApp should set up a Logger with no Loggers") { - TestLoggingWithConfigureApp.main(Array.empty) - assert(TestLoggingApp.loggerFactories.size == 1) - assert(TestLoggingWithConfigureApp.loggerFactories.head.handlers.size == 2) - // Logger is not configured with any handler/logger - assert(Logger.iterator.isEmpty) - } - - override protected def afterAll(): Unit = { - Logger.reset() - } - - override protected def beforeEach(): Unit = { - Logger.reset() - } -} diff --git a/util-logging/src/test/scala/com/twitter/logging/PolicyTest.scala b/util-logging/src/test/scala/com/twitter/logging/PolicyTest.scala index 389ef61030..ec64873ed3 100644 --- a/util-logging/src/test/scala/com/twitter/logging/PolicyTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/PolicyTest.scala @@ -1,13 +1,9 @@ package com.twitter.logging -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner - import com.twitter.util.StorageUnit +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class PolicyTest extends FunSuite { +class PolicyTest extends AnyFunSuite { import Policy._ test("Policy.parse: never") { diff --git a/util-logging/src/test/scala/com/twitter/logging/QueueingHandlerTest.scala b/util-logging/src/test/scala/com/twitter/logging/QueueingHandlerTest.scala index f4e8ae86ac..77e8326f32 100644 --- a/util-logging/src/test/scala/com/twitter/logging/QueueingHandlerTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/QueueingHandlerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -18,13 +18,11 @@ package com.twitter.logging import com.twitter.util.Local import java.util.{logging => javalog} -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.concurrent.{IntegrationPatience, Eventually} -import org.scalatest.junit.JUnitRunner +import org.scalatest.concurrent.IntegrationPatience +import org.scalatest.concurrent.Eventually +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class QueueingHandlerTest extends WordSpec with Eventually with IntegrationPatience { +class QueueingHandlerTest extends AnyFunSuite with Eventually with IntegrationPatience { class MockHandler extends Handler(BareFormatter, None) { def publish(record: javalog.LogRecord) = () @@ -40,159 +38,157 @@ class QueueingHandlerTest extends WordSpec with Eventually with IntegrationPatie logger } - "QueueingHandler" should { - "publish" in { - val logger = freshLogger() - val stringHandler = new StringHandler(BareFormatter, Some(Logger.INFO)) - val queueHandler = new QueueingHandler(stringHandler) - logger.addHandler(queueHandler) + test("QueueingHandler should publish") { + val logger = freshLogger() + val stringHandler = new StringHandler(BareFormatter, Some(Logger.INFO)) + val queueHandler = new QueueingHandler(stringHandler) + logger.addHandler(queueHandler) - logger.warning("oh noes!") - eventually { - // let thread log - assert(stringHandler.get == "oh noes!\n") - } + logger.warning("oh noes!") + eventually { + // let thread log + assert(stringHandler.get == "oh noes!\n") } + } - "publish with local context" in { - val logger = freshLogger() - val local = new Local[String] - val formatter = new Formatter { - override def format(record: javalog.LogRecord) = - local().getOrElse("MISSING!") + ":" + formatText(record) + lineTerminator - } - val stringHandler = new StringHandler(formatter, Some(Logger.INFO)) - logger.addHandler(new QueueingHandler(stringHandler)) + test("QueueingHandler should publish with local context") { + val logger = freshLogger() + val local = new Local[String] + val formatter = new Formatter { + override def format(record: javalog.LogRecord) = + local().getOrElse("MISSING!") + ":" + formatText(record) + lineTerminator + } + val stringHandler = new StringHandler(formatter, Some(Logger.INFO)) + logger.addHandler(new QueueingHandler(stringHandler)) - local() = "foo" - logger.warning("oh noes!") + local() = "foo" + logger.warning("oh noes!") - eventually { - assert(stringHandler.get == "foo:oh noes!\n") - } + eventually { + assert(stringHandler.get == "foo:oh noes!\n") } + } - "publish, drop on overflow" in { - val logger = freshLogger() - val blockingHandler = new MockHandler { - override def publish(record: javalog.LogRecord): Unit = { - Thread.sleep(100000) - } + test("QueueingHandler should publish, drop on overflow") { + val logger = freshLogger() + val blockingHandler = new MockHandler { + override def publish(record: javalog.LogRecord): Unit = { + Thread.sleep(100000) } - @volatile var droppedCount = 0 + } + @volatile var droppedCount = 0 - val queueHandler = new QueueingHandler(blockingHandler, 1) { - override protected def onOverflow(record: javalog.LogRecord): Unit = { - droppedCount += 1 - } + val queueHandler = new QueueingHandler(blockingHandler, 1) { + override protected def onOverflow(record: javalog.LogRecord): Unit = { + droppedCount += 1 } - logger.addHandler(queueHandler) + } + logger.addHandler(queueHandler) - logger.warning("1") - logger.warning("2") - logger.warning("3") + logger.warning("1") + logger.warning("2") + logger.warning("3") - eventually { - // let thread log and block - assert(droppedCount >= 1) // either 1 or 2, depending on race - } + eventually { + // let thread log and block + assert(droppedCount >= 1) // either 1 or 2, depending on race } + } - "flush" in { - val logger = freshLogger() - var wasFlushed = false - val handler = new MockHandler { - override def flush(): Unit = { wasFlushed = true } - } - val queueHandler = new QueueingHandler(handler, 1) - logger.addHandler(queueHandler) - - // smoke test: thread might write it, flush might write it - logger.warning("oh noes!") - queueHandler.flush() - assert(wasFlushed == true) + test("QueueingHandler should flush") { + val logger = freshLogger() + var wasFlushed = false + val handler = new MockHandler { + override def flush(): Unit = { wasFlushed = true } } + val queueHandler = new QueueingHandler(handler, 1) + logger.addHandler(queueHandler) - "close" in { - val logger = freshLogger() - var wasClosed = false - val handler = new MockHandler { - override def close(): Unit = { wasClosed = true } - } - val queueHandler = new QueueingHandler(handler) - logger.addHandler(queueHandler) + // smoke test: thread might write it, flush might write it + logger.warning("oh noes!") + queueHandler.flush() + assert(wasFlushed == true) + } - logger.warning("oh noes!") - eventually { - // let thread log - queueHandler.close() - assert(wasClosed == true) - } + test("QueueingHandler should close") { + val logger = freshLogger() + var wasClosed = false + val handler = new MockHandler { + override def close(): Unit = { wasClosed = true } + } + val queueHandler = new QueueingHandler(handler) + logger.addHandler(queueHandler) + + logger.warning("oh noes!") + eventually { + // let thread log + queueHandler.close() + assert(wasClosed == true) } + } - "handle exceptions in the underlying handler" in { - val logger = freshLogger() - @volatile var mustError = true - @volatile var didLog = false - val handler = new MockHandler { - override def publish(record: javalog.LogRecord): Unit = { - if (mustError) { - mustError = false - throw new Exception("Unable to log for whatever reason.") - } else { - didLog = true - } + test("QueueingHandler should handle exceptions in the underlying handler") { + val logger = freshLogger() + @volatile var mustError = true + @volatile var didLog = false + val handler = new MockHandler { + override def publish(record: javalog.LogRecord): Unit = { + if (mustError) { + mustError = false + throw new Exception("Unable to log for whatever reason.") + } else { + didLog = true } } + } - val queueHandler = new QueueingHandler(handler) - logger.addHandler(queueHandler) + val queueHandler = new QueueingHandler(handler) + logger.addHandler(queueHandler) - logger.info("fizz") - logger.info("buzz") + logger.info("fizz") + logger.info("buzz") - eventually { - // let thread log - assert(didLog == true) - } + eventually { + // let thread log + assert(didLog == true) } + } - "forward formatter to the underlying handler" in { - val handler = new MockHandler {} + test("QueueingHandler should forward formatter to the underlying handler") { + val handler = new MockHandler {} - val queueHandler = new QueueingHandler(handler) - val formatter = new Formatter() + val queueHandler = new QueueingHandler(handler) + val formatter = new Formatter() - queueHandler.setFormatter(formatter) + queueHandler.setFormatter(formatter) - assert(handler.getFormatter eq formatter) - } + assert(handler.getFormatter eq formatter) + } - "obey flag to force or not force class and method name inference" in { - val formatter = new Formatter() { - override def format(record: javalog.LogRecord) = { - val prefix = - if (record.getSourceClassName != null) record.getSourceClassName + ": " - else "" + test("QueueingHandler should obey flag to force or not force class and method name inference") { + val formatter = new Formatter() { + override def format(record: javalog.LogRecord) = { + val prefix = + if (record.getSourceClassName != null) record.getSourceClassName + ": " + else "" - prefix + formatText(record) + lineTerminator - } + prefix + formatText(record) + lineTerminator } - for (infer <- Seq(true, false)) { - val logger = freshLogger() - val stringHandler = new StringHandler(formatter, Some(Logger.INFO)) - val queueHandler = new QueueingHandler(stringHandler, Int.MaxValue, infer) - logger.addHandler(queueHandler) - val helper = new QueueingHandlerTestHelper() - helper.logSomething(logger) - - val expectedMessagePrefix = if (infer) helper.getClass.getName + ": " else "" - val expectedMessage = expectedMessagePrefix + helper.message + "\n" - - eventually { - // let thread log - assert(stringHandler.get == expectedMessage) - } + } + for (infer <- Seq(true, false)) { + val logger = freshLogger() + val stringHandler = new StringHandler(formatter, Some(Logger.INFO)) + val queueHandler = new QueueingHandler(stringHandler, Int.MaxValue, infer) + logger.addHandler(queueHandler) + val helper = new QueueingHandlerTestHelper() + helper.logSomething(logger) + + val expectedMessagePrefix = if (infer) helper.getClass.getName + ": " else "" + val expectedMessage = expectedMessagePrefix + helper.message + "\n" + + eventually { + // let thread log + assert(stringHandler.get == expectedMessage) } } } diff --git a/util-logging/src/test/scala/com/twitter/logging/ScribeHandlerTest.scala b/util-logging/src/test/scala/com/twitter/logging/ScribeHandlerTest.scala deleted file mode 100644 index 94b6910641..0000000000 --- a/util-logging/src/test/scala/com/twitter/logging/ScribeHandlerTest.scala +++ /dev/null @@ -1,259 +0,0 @@ -/* - * Copyright 2010 Twitter, Inc. - * - * 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 com.twitter -package logging - -import com.twitter.conversions.string._ -import com.twitter.conversions.time._ -import com.twitter.finagle.stats.InMemoryStatsReceiver -import com.twitter.util.{RandomSocket, Time} -import java.net.Socket -import java.util.concurrent.atomic.AtomicInteger -import java.util.concurrent.{ - ArrayBlockingQueue, - RejectedExecutionHandler, - TimeUnit, - ThreadPoolExecutor -} -import java.util.{logging => javalog} -import org.junit.runner.RunWith -import org.mockito.Mockito._ -import org.scalatest.junit.JUnitRunner -import org.scalatest.{BeforeAndAfter, WordSpec} -import org.scalatest.concurrent.Eventually -import org.scalatest.mockito.MockitoSugar -import org.scalatest.time.{Millis, Minutes, Span} - -@RunWith(classOf[JUnitRunner]) -class ScribeHandlerTest extends WordSpec with BeforeAndAfter with Eventually with MockitoSugar { - val record1 = new javalog.LogRecord(Level.INFO, "This is a message.") - record1.setMillis(1206769996722L) - record1.setLoggerName("hello") - val record2 = new javalog.LogRecord(Level.INFO, "This is another message.") - record2.setLoggerName("hello") - record2.setMillis(1206769996722L) - - // Ensure we get a free port. This assumes this port won't be reallocated - // while test is running even though SO_REUSEADDR is true - val portWithoutListener = RandomSocket.nextPort() - - "ScribeHandler" should { - - "build a scribe RPC call" in { - Time.withCurrentTimeFrozen { _ => - val scribe = ScribeHandler( - // Hack to make sure that the buffer doesn't get flushed. - port = portWithoutListener, - bufferTime = 100.milliseconds, - maxMessagesToBuffer = 10000, - formatter = new Formatter(timezone = Some("UTC")), - category = "test", - level = Some(Level.DEBUG) - ).apply() - - scribe.updateLastTransmission() - scribe.publish(record1) - scribe.publish(record2) - - assert(scribe.queue.size == 2) - assert( - scribe.makeBuffer(2).array.hexlify == ("000000b080010001000000034c6f67000000000f0001" + - "0c000000020b000100000004746573740b0002000000" + - "36494e46205b32303038303332392d30353a35333a31" + - "362e3732325d2068656c6c6f3a205468697320697320" + - "61206d6573736167652e0a000b000100000004746573" + - "740b00020000003c494e46205b32303038303332392d" + - "30353a35333a31362e3732325d2068656c6c6f3a2054" + - "68697320697320616e6f74686572206d657373616765" + - "2e0a0000") - ) - } - } - - "be able to log binary data" in { - Time.withCurrentTimeFrozen { _ => - val scribe = ScribeHandler( - port = portWithoutListener, - bufferTime = 100.milliseconds, - maxMessagesToBuffer = 10000, - formatter = new Formatter(timezone = Some("UTC")), - category = "test", - level = Some(Level.DEBUG) - ).apply() - - val bytes = Array[Byte](1, 2, 3, 4, 5) - - scribe.updateLastTransmission() - scribe.publish(bytes) - - assert(scribe.queue.peek() == bytes) - } - } - - "throw away log messages if scribe is too busy" in { - val statsReceiver = new InMemoryStatsReceiver - - val scribe = ScribeHandler( - port = portWithoutListener, - bufferTime = 5.seconds, - maxMessagesToBuffer = 1, - formatter = BareFormatter, - category = "test", - statsReceiver = statsReceiver - ).apply() - - scribe.updateLastTransmission() - scribe.publish(record1) - scribe.publish(record2) - - assert(statsReceiver.counter("dropped_records")() == 1l) - assert(statsReceiver.counter("sent_records")() == 0l) - } - - "have backoff on connection errors" in { - val statsReceiver = new InMemoryStatsReceiver - - // Use a mock ThreadPoolExecutor to force syncronize execution, - // to ensure reliably test connection backoff. - class MockThreadPoolExecutor - extends ThreadPoolExecutor( - 1, - 1, - 0L, - TimeUnit.MILLISECONDS, - new ArrayBlockingQueue[Runnable](5) - ) { - override def execute(command: Runnable): Unit = command.run() - } - - // In a loaded test environment, we can't reliably assume the port is not used, - // although the test just checked a short moment ago. - // Solution: try multiple times until it gets a connection failure. - def scribeWithConnectionFailure(retries: Int): ScribeHandler = { - val scribe = new ScribeHandler( - hostname = "localhost", - port = portWithoutListener, - category = "test", - bufferTime = 5.seconds, - connectBackoff = 15.seconds, - maxMessagesPerTransaction = 1, - maxMessagesToBuffer = 4, - formatter = BareFormatter, - level = None, - statsReceiver = statsReceiver - ) { - override private[logging] val flusher = new MockThreadPoolExecutor - } - scribe.publish(record1) - if (retries > 0 && statsReceiver.counter("connection_failed")() < 1l) - scribeWithConnectionFailure(retries - 1) - else - scribe - } - val scribe = scribeWithConnectionFailure(2) - - assert(statsReceiver.counter("connection_failed")() == 1l) - - scribe.publish(record2) - scribe.publish(record1) - scribe.publish(record2) - - scribe.flusher.shutdown() - scribe.flusher.awaitTermination(5, TimeUnit.SECONDS) - - assert(statsReceiver.counter("connection_skipped")() == 3l) - } - - // TODO rewrite deterministically when we rewrite ScribeHandler - "avoid unbounded flush queue growth" in { - val statsReceiver = new InMemoryStatsReceiver - - val scribe = ScribeHandler( - port = portWithoutListener, - bufferTime = 5.seconds, - maxMessagesToBuffer = 1, - formatter = BareFormatter, - category = "test", - statsReceiver = statsReceiver - ).apply() - - // Set up a rejectedExecutionHandler to count number of rejected tasks. - val rejected = new AtomicInteger(0) - scribe.flusher.setRejectedExecutionHandler(new RejectedExecutionHandler { - val inner = scribe.flusher.getRejectedExecutionHandler() - def rejectedExecution(r: Runnable, executor: ThreadPoolExecutor): Unit = { - rejected.getAndIncrement() - inner.rejectedExecution(r, executor) - } - }) - - // crude form to allow all 100 submission and get predictable dropping of tasks - scribe.synchronized { - scribe.publish(record1) - // wait till the ThreadPoolExecutor dequeue the first tast. - eventually(timeout(Span(2, Minutes)), interval(Span(100, Millis))) { - assert(scribe.flusher.getQueue().size() == 0) - } - // publish the rest. - for (i <- 1 until 100) { - scribe.publish(record1) - } - } - scribe.flusher.shutdown() - scribe.flusher.awaitTermination(15, TimeUnit.SECONDS) - - assert(statsReceiver.counter("published")() == 100l) - // 100 publish requests, each creates a flush task enqueued to ThreadPoolExecutor - // There will be 5 tasks stay in the threadpool's task queue. - // There will be 1 task dequeued by the executor. - // That leaves either 94 rejected tasks. - assert(rejected.get() == 94) - } - - "handle batch write errors" in { - val statsReceiver = new InMemoryStatsReceiver - val mockSocket = mock[Socket] - when(mockSocket.getOutputStream).thenReturn(null) - - val scribe = ScribeHandler( - port = portWithoutListener, - category = "test", - bufferTime = 5.seconds, - formatter = BareFormatter, - level = None, - statsReceiver = statsReceiver - ).apply() - - scribe.setSocket(Some(mockSocket)) - - // Make sure multiple records get flushed together - Time.withCurrentTimeFrozen { _ => - scribe.updateLastTransmission() - scribe.publish(record2) - scribe.publish(record1) - } - - // Throws exception writing to null output stream - scribe.flush() - - scribe.flusher.shutdown() - scribe.flusher.awaitTermination(15, TimeUnit.SECONDS) - - assert(statsReceiver.counter("dropped_records")() == 2L) - } - } -} diff --git a/util-logging/src/test/scala/com/twitter/logging/SyslogHandlerTest.scala b/util-logging/src/test/scala/com/twitter/logging/SyslogHandlerTest.scala index 920379c553..655ee18b20 100644 --- a/util-logging/src/test/scala/com/twitter/logging/SyslogHandlerTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/SyslogHandlerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,15 +16,13 @@ package com.twitter.logging -import java.net.{DatagramPacket, DatagramSocket} +import java.net.DatagramPacket +import java.net.DatagramSocket import java.util.{logging => javalog} -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class SyslogHandlerTest extends WordSpec { +class SyslogHandlerTest extends AnyFunSuite { val record1 = new javalog.LogRecord(Level.FATAL, "fatal message!") record1.setLoggerName("net.lag.whiskey.Train") record1.setMillis(1206769996722L) @@ -32,78 +30,88 @@ class SyslogHandlerTest extends WordSpec { record2.setLoggerName("net.lag.whiskey.Train") record2.setMillis(1206769996722L) - "SyslogHandler" should { - "write syslog entries" in { - // start up new syslog listener - val serverSocket = new DatagramSocket - val serverPort = serverSocket.getLocalPort + test("SyslogHandler should write syslog entries") { + // start up new syslog listener + val serverSocket = new DatagramSocket + val serverPort = serverSocket.getLocalPort - var syslog = SyslogHandler( - port = serverPort, - formatter = new SyslogFormatter( - timezone = Some("UTC"), - hostname = "raccoon.local" - ) - ).apply() - syslog.publish(record1) - syslog.publish(record2) - - SyslogFuture.sync - val p = new DatagramPacket(new Array[Byte](1024), 1024) - serverSocket.receive(p) - assert( - new String(p.getData, 0, p.getLength) == "<9>2008-03-29T05:53:16 raccoon.local whiskey: fatal message!" - ) - serverSocket.receive(p) - assert( - new String(p.getData, 0, p.getLength) == "<11>2008-03-29T05:53:16 raccoon.local whiskey: error message!" + var syslog = SyslogHandler( + port = serverPort, + formatter = new SyslogFormatter( + timezone = Some("UTC"), + hostname = "raccoon.local" ) - } + ).apply() + syslog.publish(record1) + syslog.publish(record2) - "with server name" in { - // start up new syslog listener - val serverSocket = new DatagramSocket - val serverPort = serverSocket.getLocalPort + SyslogFuture.sync() + val p = new DatagramPacket(new Array[Byte](1024), 1024) + serverSocket.receive(p) + assert( + new String( + p.getData, + 0, + p.getLength) == "<9>2008-03-29T05:53:16 raccoon.local whiskey: fatal message!" + ) + serverSocket.receive(p) + assert( + new String( + p.getData, + 0, + p.getLength) == "<11>2008-03-29T05:53:16 raccoon.local whiskey: error message!" + ) + } - var syslog = SyslogHandler( - port = serverPort, - formatter = new SyslogFormatter( - serverName = Some("pingd"), - timezone = Some("UTC"), - hostname = "raccoon.local" - ) - ).apply() - syslog.publish(record1) + test("SyslogHandler with with server name") { + // start up new syslog listener + val serverSocket = new DatagramSocket + val serverPort = serverSocket.getLocalPort - SyslogFuture.sync - val p = new DatagramPacket(new Array[Byte](1024), 1024) - serverSocket.receive(p) - assert( - new String(p.getData, 0, p.getLength) == "<9>2008-03-29T05:53:16 raccoon.local [pingd] whiskey: fatal message!" + var syslog = SyslogHandler( + port = serverPort, + formatter = new SyslogFormatter( + serverName = Some("pingd"), + timezone = Some("UTC"), + hostname = "raccoon.local" ) - } + ).apply() + syslog.publish(record1) - "with BSD time format" in { - // start up new syslog listener - val serverSocket = new DatagramSocket - val serverPort = serverSocket.getLocalPort + SyslogFuture.sync() + val p = new DatagramPacket(new Array[Byte](1024), 1024) + serverSocket.receive(p) + assert( + new String( + p.getData, + 0, + p.getLength) == "<9>2008-03-29T05:53:16 raccoon.local [pingd] whiskey: fatal message!" + ) + } - var syslog = SyslogHandler( - port = serverPort, - formatter = new SyslogFormatter( - useIsoDateFormat = false, - timezone = Some("UTC"), - hostname = "raccoon.local" - ) - ).apply() - syslog.publish(record1) + test("SyslogHandler with BSD time format") { + // start up new syslog listener + val serverSocket = new DatagramSocket + val serverPort = serverSocket.getLocalPort - SyslogFuture.sync - val p = new DatagramPacket(new Array[Byte](1024), 1024) - serverSocket.receive(p) - assert( - new String(p.getData, 0, p.getLength) == "<9>Mar 29 05:53:16 raccoon.local whiskey: fatal message!" + var syslog = SyslogHandler( + port = serverPort, + formatter = new SyslogFormatter( + useIsoDateFormat = false, + timezone = Some("UTC"), + hostname = "raccoon.local" ) - } + ).apply() + syslog.publish(record1) + + SyslogFuture.sync() + val p = new DatagramPacket(new Array[Byte](1024), 1024) + serverSocket.receive(p) + assert( + new String( + p.getData, + 0, + p.getLength) == "<9>Mar 29 05:53:16 raccoon.local whiskey: fatal message!" + ) } } diff --git a/util-logging/src/test/scala/com/twitter/logging/ThrottledHandlerTest.scala b/util-logging/src/test/scala/com/twitter/logging/ThrottledHandlerTest.scala index e2f0b65de8..b5fc3dfde7 100644 --- a/util-logging/src/test/scala/com/twitter/logging/ThrottledHandlerTest.scala +++ b/util-logging/src/test/scala/com/twitter/logging/ThrottledHandlerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,80 +16,77 @@ package com.twitter.logging -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.{BeforeAndAfter, WordSpec} +import com.twitter.conversions.DurationOps._ +import com.twitter.io.TempFolder +import com.twitter.util.Time +import org.scalatest.BeforeAndAfter +import org.scalatest.funsuite.AnyFunSuite -import com.twitter.conversions.time._ -import com.twitter.util.{TempFolder, Time} -@RunWith(classOf[JUnitRunner]) -class ThrottledHandlerTest extends WordSpec with BeforeAndAfter with TempFolder { - private var handler: StringHandler = null +class ThrottledHandlerTest extends AnyFunSuite with BeforeAndAfter with TempFolder { + private var handler: StringHandler = _ - "ThrottledHandler" should { - before { - Logger.clearHandlers - handler = new StringHandler(BareFormatter, None) - } + before { + Logger.clearHandlers() + handler = new StringHandler(BareFormatter, None) + } - after { - Logger.clearHandlers - } + after { + Logger.clearHandlers() + } - "throttle keyed log messages" in { - val log = Logger() - val throttledLog = new ThrottledHandler(handler, 1.second, 3) - Time.withCurrentTimeFrozen { timeCtrl => - log.addHandler(throttledLog) + test("ThrottledHandler should throttle keyed log messages") { + val log = Logger() + val throttledLog = new ThrottledHandler(handler, 1.second, 3) + Time.withCurrentTimeFrozen { timeCtrl => + log.addHandler(throttledLog) - log.error("apple: %s", "help!") - log.error("apple: %s", "help 2!") - log.error("orange: %s", "orange!") - log.error("orange: %s", "orange!") - log.error("apple: %s", "help 3!") - log.error("apple: %s", "help 4!") - log.error("apple: %s", "help 5!") - timeCtrl.advance(2.seconds) - log.error("apple: %s", "done.") + log.error("apple: %s", "help!") + log.error("apple: %s", "help 2!") + log.error("orange: %s", "orange!") + log.error("orange: %s", "orange!") + log.error("apple: %s", "help 3!") + log.error("apple: %s", "help 4!") + log.error("apple: %s", "help 5!") + timeCtrl.advance(2.seconds) + log.error("apple: %s", "done.") - assert( - handler.get.split("\n").toList == List( - "apple: help!", - "apple: help 2!", - "orange: orange!", - "orange: orange!", - "apple: help 3!", - "(swallowed 2 repeating messages)", - "apple: done." - ) + assert( + handler.get.split("\n").toList == List( + "apple: help!", + "apple: help 2!", + "orange: orange!", + "orange: orange!", + "apple: help 3!", + "(swallowed 2 repeating messages)", + "apple: done." ) - } + ) } + } - "log the summary even if nothing more is logged with that name" in { - Time.withCurrentTimeFrozen { time => - val log = Logger() - val throttledLog = new ThrottledHandler(handler, 1.second, 3) - log.addHandler(throttledLog) - log.error("apple: %s", "help!") - log.error("apple: %s", "help!") - log.error("apple: %s", "help!") - log.error("apple: %s", "help!") - log.error("apple: %s", "help!") + test("ThrottledHandler should log the summary even if nothing more is logged with that name") { + Time.withCurrentTimeFrozen { time => + val log = Logger() + val throttledLog = new ThrottledHandler(handler, 1.second, 3) + log.addHandler(throttledLog) + log.error("apple: %s", "help!") + log.error("apple: %s", "help!") + log.error("apple: %s", "help!") + log.error("apple: %s", "help!") + log.error("apple: %s", "help!") - time.advance(2.seconds) - log.error("hello.") + time.advance(2.seconds) + log.error("hello.") - assert( - handler.get.split("\n").toList == List( - "apple: help!", - "apple: help!", - "apple: help!", - "(swallowed 2 repeating messages)", - "hello." - ) + assert( + handler.get.split("\n").toList == List( + "apple: help!", + "apple: help!", + "apple: help!", + "(swallowed 2 repeating messages)", + "hello." ) - } + ) } } } diff --git a/util-mock/PROJECT b/util-mock/PROJECT new file mode 100644 index 0000000000..0f66d02b27 --- /dev/null +++ b/util-mock/PROJECT @@ -0,0 +1,5 @@ +owners: + - jillianc + - mbezoyan +watchers: + - csl-team:ldap diff --git a/util-mock/src/main/scala/com/twitter/util/mock/BUILD b/util-mock/src/main/scala/com/twitter/util/mock/BUILD new file mode 100644 index 0000000000..ede98b769d --- /dev/null +++ b/util-mock/src/main/scala/com/twitter/util/mock/BUILD @@ -0,0 +1,19 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-mockito", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/mockito:mockito-scala", + ], + exports = [ + "3rdparty/jvm/org/mockito:mockito-scala", + ], +) diff --git a/util-mock/src/main/scala/com/twitter/util/mock/Mockito.scala b/util-mock/src/main/scala/com/twitter/util/mock/Mockito.scala new file mode 100644 index 0000000000..9e87fe1dfa --- /dev/null +++ b/util-mock/src/main/scala/com/twitter/util/mock/Mockito.scala @@ -0,0 +1,167 @@ +package com.twitter.util.mock + +import org.mockito.{ArgumentMatchersSugar, IdiomaticMockito} + +/** + * Helper for Mockito Scala sugar with [[https://github.com/mockito/mockito-scala#idiomatic-mockito idiomatic stubbing]]. + * Java users are encouraged to use `org.mockito.Mockito` directly. + * + * Note that the Specs2 `smartMock[]` or `mock[].smart` is the default behavior + * for [[https://github.com/mockito/mockito-scala Mockito Scala]]. + * + * =Usage= + * + * This trait uses `org.mockito.IdiomaticMockito` which is heavily influenced by ScalaTest Matchers. + * + * To use, mix in the [[com.twitter.util.mock.Mockito]] trait where desired. + * + * ==Create a new mock== + * + * {{{ + * + * trait Foo { + * def bar: String + * def bar(v: Int): Int + * } + * + * class MyTest extends AnyFunSuite with Mockito { + * val aMock = mock[Foo] + * } + * }}} + * + * ==Expect behavior== + * + * {{{ + * // "when" equivalents + * aMock.bar returns "mocked!" + * aMock.bar returns "mocked!" andThen "mocked again!" + * aMock.bar shouldCall realMethod + * aMock.bar.shouldThrow[IllegalArgumentException] + * aMock.bar throws new IllegalArgumentException + * aMock.bar answers "mocked!" + * aMock.bar(*) answers ((i: Int) => i * 10) + * + * // "do-when" equivalents + * "mocked!" willBe returned by aMock.bar + * "mocked!" willBe answered by aMock.bar + * ((i: Int) => i * 10) willBe answered by aMock.bar(*) + * theRealMethod willBe called by aMock.bar + * new IllegalArgumentException willBe thrown by aMock.bar + * aMock.bar.doesNothing() // doNothing().when(aMock).bar + * + * // verifications + * aMock wasNever called // verifyZeroInteractions(aMock) + * aMock.bar was called + * aMock.bar(*) was called // '*' is shorthand for 'any()' or 'any[T]' + * aMock.bar(any[Int]) was called // same as above but with typed input matcher + * + * aMock.bar wasCalled onlyHere + * aMock.bar wasNever called + * + * aMock.bar wasCalled twice + * aMock.bar wasCalled 2.times + * + * aMock.bar wasCalled fourTimes + * aMock.bar wasCalled 4.times + * + * aMock.bar wasCalled atLeastFiveTimes + * aMock.bar wasCalled atLeast(fiveTimes) + * aMock.bar wasCalled atLeast(5.times) + * + * aMock.bar wasCalled atMostSixTimes + * aMock.bar wasCalled atMost(sixTimes) + * aMock.bar wasCalled atMost(6.times) + * + * aMock.bar wasCalled (atLeastSixTimes within 2.seconds) // verify(aMock, timeout(2000).atLeast(6)).bar + * + * aMock wasNever calledAgain // verifyNoMoreInteractions(aMock) + * + * InOrder(mock1, mock2) { implicit order => + * mock2.someMethod() was called + * mock1.anotherMethod() was called + * } + * }}} + * + * Note the 'dead code' warning that can happen when using 'any' or '*' + * [[https://github.com/mockito/mockito-scala#dead-code-warning matchers]]. + * + * ==Mixing and matching matchers== + * + * Using the idiomatic syntax also allows for mixing argument matchers with real values. E.g., you + * are no longer forced to use argument matchers for all parameters as soon as you use one. E.g., + * + * {{{ + * trait Foo { + * def bar(v: Int, v2: Int, v3: Int = 42): Int + * } + * + * class MyTest extends AnyFunSuite with Mockito { + * val aMock = mock[Foo] + * + * aMock.bar(1, 2) returns "mocked!" + * aMock.bar(1, *) returns "mocked!" + * aMock.bar(1, any[Int]) returns "mocked!" + * aMock.bar(*, *) returns "mocked!" + * aMock.bar(any[Int], any[Int]) returns "mocked!" + * aMock.bar(*, *, 3) returns "mocked!" + * aMock.bar(any[Int], any[Int], 3) returns "mocked!" + * + * "mocked!" willBe returned by aMock.bar(1, 2) + * "mocked!" willBe returned by aMock.bar(1, *) + * "mocked!" willBe returned by aMock.bar(1, any[Int]) + * "mocked!" willBe returned by aMock.bar(*, *) + * "mocked!" willBe returned by aMock.bar(any[Int], any[Int]) + * "mocked!" willBe returned by aMock.bar(*, *, 3) + * "mocked!" willBe returned by aMock.bar(any[Int], any[Int], 3) + * + * aMock.bar(1, 2) was called + * aMock.bar(1, *) was called + * aMock.bar(1, any[Int]) was called + * aMock.bar(*, *) was called + * aMock.bar(any[Int], any[Int]) was called + * aMock.bar(*, *, 3) was called + * aMock.bar(any[Int], any[Int], 3) was called + * } + * }}} + * + * See [[https://github.com/mockito/mockito-scala#mixing-normal-values-with-argument-matchers Mix-and-Match]] + * for more information including a caveat around curried functions with default arguments. + * + * ==Numeric Matchers== + * + * Numeric comparisons are possible for argument matching, e.g., + * + * {{{ + * aMock.method(5) + * + * aMock.method(n > 4.99) was called + * aMock.method(n >= 5) was called + * aMock.method(n < 5.1) was called + * aMock.method(n <= 5) was called + * }}} + * + * See [[https://github.com/mockito/mockito-scala#numeric-matchers Numeric Matchers]]. + * + * ==Vargargs== + * + * Most matches will deal with varargs out of the box, just note when using the 'eqTo' matcher to + * apply it to all the arguments as one (not individually). + * + * See [[https://github.com/mockito/mockito-scala#varargs Varargs]]. + * + * ==More Information== + * + * See the [[https://github.com/mockito/mockito-scala#idiomatic-mockito IdiomaticMockito]] documentation + * for more specific information and the [[https://github.com/mockito/mockito-scala Mockito Scala]] + * [[https://github.com/mockito/mockito-scala#getting-started Getting Started]] documentation for general + * information. + * + * see `org.mockito.IdiomaticMockito` + * see `org.mockito.ArgumentMatchersSugar` + */ +trait Mockito extends IdiomaticMockito with ArgumentMatchersSugar + +/** + * Simple object to allow the usage of the trait without mixing it in + */ +object Mockito extends Mockito diff --git a/util-mock/src/test/java/com/twitter/util/mock/BUILD.bazel b/util-mock/src/test/java/com/twitter/util/mock/BUILD.bazel new file mode 100644 index 0000000000..22faa8e0fc --- /dev/null +++ b/util-mock/src/test/java/com/twitter/util/mock/BUILD.bazel @@ -0,0 +1,9 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + ], +) diff --git a/util-mock/src/test/java/com/twitter/util/mock/JavaFoo.java b/util-mock/src/test/java/com/twitter/util/mock/JavaFoo.java new file mode 100644 index 0000000000..012851460d --- /dev/null +++ b/util-mock/src/test/java/com/twitter/util/mock/JavaFoo.java @@ -0,0 +1,5 @@ +package com.twitter.util.mock; + +public interface JavaFoo { + Integer varargMethod(Integer... arg); +} diff --git a/util-mock/src/test/scala/com/twitter/util/mock/BUILD.bazel b/util-mock/src/test/scala/com/twitter/util/mock/BUILD.bazel new file mode 100644 index 0000000000..9680026e4c --- /dev/null +++ b/util-mock/src/test/scala/com/twitter/util/mock/BUILD.bazel @@ -0,0 +1,20 @@ +junit_tests( + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/mockito:mockito-scala", + "3rdparty/jvm/org/mockito:mockito-scala-scalatest", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:scala", + "util/util-mock/src/main/scala/com/twitter/util/mock", + "util/util-mock/src/test/java/com/twitter/util/mock", + ], +) diff --git a/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitTest.scala b/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitTest.scala new file mode 100644 index 0000000000..1a37c2b551 --- /dev/null +++ b/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitTest.scala @@ -0,0 +1,751 @@ +package com.twitter.util.mock + +import com.twitter.conversions.DurationOps._ +import com.twitter.util._ +import java.io.{File, FileOutputStream, ObjectOutputStream} +import java.util.concurrent.atomic.AtomicInteger +import org.junit.runner.RunWith +import org.mockito.MockitoSugar +import org.mockito.captor.ArgCaptor +import org.mockito.exceptions.verification.WantedButNotInvoked +import org.mockito.invocation.InvocationOnMock +import org.mockito.stubbing.{CallsRealMethods, DefaultAnswer, ScalaFirstStubbing} +import org.scalatest.concurrent.PatienceConfiguration.{Interval, Timeout} +import org.scalatest.concurrent.ScalaFutures.{FutureConcept, PatienceConfig, whenReady} +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatest.time.{Millis, Seconds, Span} +import org.scalatestplus.junit.JUnitRunner +import scala.annotation.varargs +import scala.language.{higherKinds, implicitConversions} +import scala.reflect.io.AbstractFile + +object MockitoTraitTest { + + /* Test classes pulled from MockitoScala: + https://github.com/mockito/mockito-scala/blob/release/1.x/scalatest/src/test/scala/user/org/mockito/TestModel.scala*/ + + class X { + def bar = "not mocked" + def iHaveByNameArgs(normal: String, byName: => String, byName2: => String): String = + "not mocked" + def iHaveSomeDefaultArguments(noDefault: String, default: String = "default value"): String = + "not mocked" + def iStartWithByNameArgs(byName: => String, normal: String): String = "not mocked" + def iHavePrimitiveByNameArgs(byName: => Int, normal: String): String = "not mocked" + def iHaveByNameAndVarArgs( + byName: => String, + normal: String, + byName2: => String, + normal2: String, + vararg: String* + )( + byName3: => String, + normal3: String, + vararg3: String* + ): String = "not mocked" + + def iHaveFunction0Args(normal: String, f0: () => String): String = "not mocked" + def iHaveByNameAndFunction0Args(normal: String, f0: () => String, byName: => String): String = + "not mocked" + def returnBar: Bar = new Bar + def doSomethingWithThisIntAndString(v: Int, v2: String): ValueCaseClassInt = ValueCaseClassInt( + v) + def returnsValueCaseClass: ValueCaseClassInt = ValueCaseClassInt(-1) + def returnsValueCaseClass(i: Int): ValueCaseClassInt = ValueCaseClassInt(i) + def baz(i: Int, b: Baz2): String = "not mocked" + } + + trait Y { + def varargMethod(s: String, arg: Int*): Int = -1 + def varargMethod(arg: Int*): Int = -1 + @varargs def javaVarargMethod(arg: Int*): Int = -1 + def byNameMethod(arg: => Int): Int = -1 + def traitMethod(arg: Int): ValueCaseClassInt = ValueCaseClassInt(arg) + def traitMethodWithDefaultArgs(defaultArg: Int = 30, anotherDefault: String = "hola"): Int = -1 + } + + class XWithY extends X with Y + + class Complex extends X + + class Foo { + def bar(a: String) = "bar" + } + + class Foo2 { + def valueClass(n: Int, v: ValueClass): String = ??? + def valueCaseClass(n: Int, v: ValueCaseClassInt): String = ??? + } + + class Bar { + def iAlsoHaveSomeDefaultArguments( + noDefault: String, + default: String = "default value" + ): String = "not mocked" + def iHaveDefaultArgs(v: String = "default"): String = "not mocked" + } + + trait Baz { + def varargMethod(s: String, arg: Int*): Int = -1 + def varargMethod(arg: Int*): Int = -1 + @varargs def javaVarargMethod(arg: Int*): Int = -1 + def byNameMethod(arg: => Int): Int = -1 + def traitMethod(arg: Int): ValueCaseClassInt = ValueCaseClassInt(arg) + def traitMethodWithDefaultArgs(defaultArg: Int = 30, anotherDefault: String = "hola"): Int = -1 + } + + implicit class ValueClass(val v: String) extends AnyVal + case class ValueCaseClassInt(v: Int) extends AnyVal + case class ValueCaseClassString(v: String) extends AnyVal + case class GenericValue[T](v: Int) extends AnyVal + + class HigherKinded[F[_]] { + def method: F[Either[String, String]] = null.asInstanceOf[F[Either[String, String]]] + def method2: F[Either[String, String]] = null.asInstanceOf[F[Either[String, String]]] + } + + class ConcreteHigherKinded extends HigherKinded[Option] + + class Implicit[T] + + class Org { + def bar = "not mocked" + def baz = "not mocked" + + def defaultParams(a: String, defaultParam1: Boolean = false, defaultParam2: Int = 1): String = + "not mocked" + def curriedDefaultParams( + a: String, + defaultParam1: Boolean = false + )( + defaultParam2: Int = a.length + ): String = "not mocked" + + def doSomethingWithThisInt(v: Int): Int = -1 + def doSomethingWithThisIntAndString(v: Int, v2: String): String = "not mocked" + def doSomethingWithThisIntAndStringAndBoolean(v: Int, v2: String, v3: Boolean): String = + "not mocked" + def returnBar: Bar = new Bar + def highOrderFunction(f: Int => String): String = "not mocked" + def iReturnAFunction(v: Int): Int => String = i => (i * v).toString + def iBlowUp(v: Int, v2: String): String = throw new IllegalArgumentException("I was called!") + def iHaveTypeParamsAndImplicits[A, B](a: A, b: B)(implicit v3: Implicit[A]): String = + "not mocked" + def valueClass(n: Int, v: ValueClass): String = "not mocked" + def valueCaseClass(n: Int, v: ValueCaseClassInt): String = "not mocked" + def returnsValueCaseClassInt: ValueCaseClassInt = ValueCaseClassInt(-1) + def returnsValueCaseClassString: ValueCaseClassString = ValueCaseClassString("not mocked") + def takesManyValueClasses( + v: ValueClass, + v1: ValueCaseClassInt, + v2: ValueCaseClassString + ): String = "not mocked" + def baz(i: Int, b: Baz): String = "not mocked" + def fooWithVarArg(bells: String*): Unit = () + def fooWithActualArray(bells: Array[String]): Unit = () + def fooWithSecondParameterList(bell: String)(awaitable: Awaitable[Int]): Unit = () + def fooWithVarArgAndSecondParameterList(bells: String*)(awaitable: Awaitable[Int]): Unit = + () + def fooWithActualArrayAndSecondParameterList( + bells: Array[String] + )( + awaitable: Awaitable[Int] + ): Unit = () + def valueClassWithVarArg(monitors: Monitor*): Unit = () + def valueClassWithSecondParameterList(monitor: Monitor)(awaitable: Awaitable[Int]): Unit = () + def valueClassWithVarArgAndSecondParameterList( + monitors: Monitor* + )( + awaitable: Awaitable[Int] + ): Unit = () + + def option: Option[String] = None + def option(a: String, b: Int): Option[String] = None + def printGenericValue(value: GenericValue[String]): String = "not mocked" + def unit(): Unit = ??? + } + + case class Baz2(param1: Int, param2: String) + + trait ParameterizedTrait[+E] { def m(): E } + class ParameterizedTraitInt extends ParameterizedTrait[Int] { def m() = -1 } +} + +@RunWith(classOf[JUnitRunner]) +class MockitoTraitTest extends AnyFunSuite with Matchers with Mockito { + import MockitoTraitTest._ + + test("mocks are called with the right arguments") { + val foo = mock[Foo] + + foo.bar(any[String]) returns "mocked" + foo.bar("alice") should equal("mocked") + foo.bar("alice") was called + } + + test("create a mock with name") { + val aMock = mock[Complex]("Named Mock") + aMock.toString should be("Named Mock") + } + + test("work with mixins") { + val aMock = mock[XWithY] + + aMock.bar returns "mocked!" + aMock.traitMethod(any[Int]) returns ValueCaseClassInt(69) + + aMock.bar should be("mocked!") + aMock.traitMethod(30) should be(ValueCaseClassInt(69)) + + MockitoSugar.verify(aMock).traitMethod(30) + } + + test("work with inline mixins") { + val aMock: X with Baz = mock[X with Baz] + + aMock.bar returns "mocked!" + aMock.traitMethod(any[Int]) returns ValueCaseClassInt(69) + aMock.varargMethod(1, 2, 3) returns 42 + aMock.javaVarargMethod(1, 2, 3) returns 42 + aMock.byNameMethod(69) returns 42 + + aMock.bar should be("mocked!") + aMock.traitMethod(30) should be(ValueCaseClassInt(69)) + aMock.varargMethod(1, 2, 3) should equal(42) + aMock.javaVarargMethod(1, 2, 3) should equal(42) + aMock.byNameMethod(69) should equal(42) + + MockitoSugar.verify(aMock).traitMethod(30) + MockitoSugar.verify(aMock).varargMethod(1, 2, 3) + MockitoSugar.verify(aMock).javaVarargMethod(1, 2, 3) + MockitoSugar.verify(aMock).byNameMethod(69) + a[WantedButNotInvoked] should be thrownBy MockitoSugar.verify(aMock).varargMethod(1, 2) + a[WantedButNotInvoked] should be thrownBy MockitoSugar.verify(aMock).javaVarargMethod(1, 2) + a[WantedButNotInvoked] should be thrownBy MockitoSugar.verify(aMock).byNameMethod(71) + } + + test("work with java varargs") { + val aMock = mock[JavaFoo] + + aMock.varargMethod(1, 2, 3) returns 42 + aMock.varargMethod(1, 2, 3) should equal(42) + + MockitoSugar.verify(aMock).varargMethod(1, 2, 3) + a[WantedButNotInvoked] should be thrownBy MockitoSugar.verify(aMock).varargMethod(1, 2) + } + + test("work when getting varargs from collections") { + val aMock = mock[Baz] + + aMock.varargMethod("Hello", List(1, 2, 3): _*) returns 42 + aMock.varargMethod("Hello", 1, 2, 3) should equal(42) + MockitoSugar.verify(aMock).varargMethod("Hello", 1, 2, 3) + } + + test("work when getting varargs from collections (with matchers)") { + val aMock = mock[Baz] + + aMock.varargMethod(eqTo("Hello"), eqTo(List(1, 2, 3)): _*) returns 42 + aMock.varargMethod("Hello", 1, 2, 3) should equal(42) + MockitoSugar.verify(aMock).varargMethod("Hello", 1, 2, 3) + } + + test("be serializable") { + val list = mock[java.util.List[String]](withSettings.name("list1").serializable()) + list.get(eqTo(3)) returns "mocked" + list.get(3) should be("mocked") + + val file = File.createTempFile("mock", "tmp") + file.deleteOnExit() + + val oos = new ObjectOutputStream(new FileOutputStream(file)) + try { + oos.writeObject(list) + } finally { + oos.close() + } + } + + test("work with type aliases") { + type MyType = String + + val aMock = mock[MyType => String] + + aMock("Hello") returns "Goodbye" + aMock("Hello") should be("Goodbye") + } + + test("not mock final methods") { + val abstractFile = mock[AbstractFile] + + abstractFile.path returns "SomeFile.scala" + abstractFile.path should be("SomeFile.scala") + } + + test("create a valid spy for lambdas and anonymous classes") { + val aSpy = spyLambda((arg: String) => s"Got: $arg") + + aSpy.apply(any[String]) returns "mocked!" + aSpy("hi!") should be("mocked!") + MockitoSugar.verify(aSpy).apply("hi!") + } + + test("eqTo syntax") { + val aMock = mock[Foo2] + + aMock.valueClass(1, eqTo(ValueClass("meh"))) returns "mocked!" + aMock.valueClass(1, ValueClass("meh")) should be("mocked!") + aMock.valueClass(1, eqTo(ValueClass("meh"))) was called + + aMock.valueClass(any[Int], ValueClass("moo")) returns "mocked!" + aMock.valueClass(11, ValueClass("moo")) should be("mocked!") + aMock.valueClass(any[Int], ValueClass("moo")) was called + + val valueClass = ValueClass("blah") + aMock.valueClass(1, eqTo(valueClass)) returns "mocked!" + aMock.valueClass(1, valueClass) should be("mocked!") + aMock.valueClass(1, eqTo(valueClass)) was called + + aMock.valueCaseClass(2, eqTo(ValueCaseClassInt(100))) returns "mocked!" + aMock.valueCaseClass(2, ValueCaseClassInt(100)) should be("mocked!") + aMock.valueCaseClass(2, eqTo(ValueCaseClassInt(100))) was called + + val caseClassValue = ValueCaseClassInt(100) + aMock.valueCaseClass(3, eqTo(caseClassValue)) returns "mocked!" + aMock.valueCaseClass(3, caseClassValue) should be("mocked!") + aMock.valueCaseClass(3, eqTo(caseClassValue)) was called + + aMock.valueCaseClass(any[Int], ValueCaseClassInt(200)) returns "mocked!" + aMock.valueCaseClass(4, ValueCaseClassInt(200)) should be("mocked!") + aMock.valueCaseClass(any[Int], ValueCaseClassInt(200)) was called + } + + test("create a mock with default answer") { + val aMock = mock[Complex](CallsRealMethods) + + aMock.bar should be("not mocked") + } + + test("create a mock with default answer from implicit scope") { + implicit val defaultAnswer: DefaultAnswer = CallsRealMethods + + val aMock = mock[Complex] + aMock.bar should be("not mocked") + } + + test("stub a value class return value") { + val aMock = mock[Complex] + + aMock.returnsValueCaseClass returns ValueCaseClassInt(100) andThen ValueCaseClassInt(200) + + aMock.returnsValueCaseClass should be(ValueCaseClassInt(100)) + aMock.returnsValueCaseClass should be(ValueCaseClassInt(200)) + } + + test("stub a map") { + val mocked = mock[Map[String, String]] + mocked(any[String]) returns "123" + mocked("key") should be("123") + } + + test("work with stubbing value type parameters") { + val aMock = mock[ParameterizedTraitInt] + + aMock.m() returns 100 + + aMock.m() should be(100) + aMock.m() was called + } + + test("deal with verifying value type parameters") { + val aMock = mock[ParameterizedTraitInt] + //this has to be done separately as the verification in the other test would return the stubbed value so the + //null problem on the primitive would not happen + aMock.m() wasNever called + } + + test("stub a spy that would fail if the real impl is called") { + val aSpy = spy(new Org) + + an[IllegalArgumentException] should be thrownBy { + aSpy.iBlowUp(any[Int], any[String]) returns "mocked!" + } + + "mocked!" willBe returned by aSpy.iBlowUp(any[Int], "ok") + + aSpy.iBlowUp(1, "ok") should be("mocked!") + aSpy.iBlowUp(2, "ok") should be("mocked!") + + an[IllegalArgumentException] should be thrownBy { + aSpy.iBlowUp(2, "not ok") + } + } + + test("work with parameterized return types") { + val aMock = mock[HigherKinded[Option]] + + aMock.method returns Some(Right("Mocked!")) + aMock.method.get.right.get should be("Mocked!") + aMock.method2 + } + + test("work with parameterized return types declared in parents") { + val aMock = mock[ConcreteHigherKinded] + + aMock.method returns Some(Right("Mocked!")) + + aMock.method.get.right.get should be("Mocked!") + aMock.method2 + } + + test("work with higher kinded types and auxiliary methods") { + def whenGetById[F[_]](algebra: HigherKinded[F]): ScalaFirstStubbing[F[Either[String, String]]] = + MockitoSugar.when(algebra.method) + + val aMock = mock[HigherKinded[Option]] + + whenGetById(aMock) thenReturn Some(Right("Mocked!")) + aMock.method.get.right.get should be("Mocked!") + } + + test("stub a spy with an answer") { + val aSpy = spy(new Org) + + ((i: Int) => i * 10 + 2) willBe answered by aSpy.doSomethingWithThisInt(any[Int]) + ((i: Int, s: String) => (i * 10 + s.toInt).toString) willBe answered by aSpy + .doSomethingWithThisIntAndString(any[Int], any[String]) + ( + ( + i: Int, + s: String, + boolean: Boolean + ) => (i * 10 + s.toInt).toString + boolean + ) willBe answered by aSpy + .doSomethingWithThisIntAndStringAndBoolean(any[Int], any[String], v3 = true) + val counter = new AtomicInteger(1) + (() => counter.getAndIncrement().toString) willBe answered by aSpy.bar + counter.getAndIncrement().toString willBe answered by aSpy.baz + + counter.get should equal(1) + aSpy.bar should be("1") + counter.get should equal(2) + aSpy.baz should be("2") + counter.get should equal(3) + + aSpy.doSomethingWithThisInt(4) should equal(42) + aSpy.doSomethingWithThisIntAndString(4, "2") should be("42") + aSpy.doSomethingWithThisIntAndStringAndBoolean(4, "2", v3 = true) should be("42true") + aSpy.doSomethingWithThisIntAndStringAndBoolean(4, "2", v3 = false) should be("not mocked") + } + + test("create a mock where with mixed matchers and normal parameters (answer)") { + val org = mock[Org] + + org.doSomethingWithThisIntAndString(*, "test") answers "mocked!" + + org.doSomethingWithThisIntAndString(3, "test") should be("mocked!") + org.doSomethingWithThisIntAndString(5, "test") should be("mocked!") + org.doSomethingWithThisIntAndString(5, "est") should not be "mocked!" + } + + test("create a mock with nice answer API (multiple params)") { + val aMock = mock[Complex] + + aMock.doSomethingWithThisIntAndString(any[Int], any[String]) answers ((i: Int, s: String) => + ValueCaseClassInt(i * 10 + s.toInt)) andThenAnswer ((i: Int, _: String) => + ValueCaseClassInt(i)) + + aMock.doSomethingWithThisIntAndString(4, "2") should be(ValueCaseClassInt(42)) + aMock.doSomethingWithThisIntAndString(4, "2") should be(ValueCaseClassInt(4)) + } + + test("simplify answer API (invocation usage)") { + val org = mock[Org] + + org.doSomethingWithThisInt(any[Int]) answers ((i: InvocationOnMock) => i.arg[Int](0) * 10 + 2) + org.doSomethingWithThisInt(4) should equal(42) + } + + test("chain answers") { + val org = mock[Org] + + org.doSomethingWithThisInt(any[Int]) answers ((i: Int) => i * 10 + 2) andThenAnswer ((i: Int) => + i * 15 + 9) + + org.doSomethingWithThisInt(4) should equal(42) + org.doSomethingWithThisInt(4) should equal(69) + } + + test("chain answers (invocation usage)") { + val org = mock[Org] + + org.doSomethingWithThisInt(any[Int]) answers ((i: InvocationOnMock) => + i.arg[Int](0) * 10 + 2) andThenAnswer ((i: InvocationOnMock) => i.arg[Int](0) * 15 + 9) + + org.doSomethingWithThisInt(4) should equal(42) + org.doSomethingWithThisInt(4) should equal(69) + } + + test("allow using less params than method on answer stubbing") { + val org = mock[Org] + + org.doSomethingWithThisIntAndStringAndBoolean(any[Int], any[String], any[Boolean]) answers ( + (i: Int, s: String) => (i * 10 + s.toInt).toString) + org.doSomethingWithThisIntAndStringAndBoolean(4, "2", v3 = true) should be("42") + } + + test("stub a mock inline that has default args") { + val aMock = mock[Org] + + aMock.returnBar returns mock[Bar] andThen mock[Bar] + + aMock.returnBar should be(a[Bar]) + aMock.returnBar should be(a[Bar]) + } + + test("stub a high order function") { + val org = mock[Org] + + org.highOrderFunction(any[Int => String]) returns "mocked!" + + org.highOrderFunction(_.toString) should be("mocked!") + } + + test("stub a method that returns a function") { + val org = mock[Org] + + org + .iReturnAFunction(any[Int]).shouldReturn(_.toString).andThen(i => + (i * 2).toString).andThenCallRealMethod() + + org.iReturnAFunction(0)(42) should be("42") + org.iReturnAFunction(0)(42) should be("84") + org.iReturnAFunction(3)(3) should be("9") + } + + //useful if we want to delay the evaluation of whatever we are returning until the method is called + test("simplify stubbing an answer where we don't care about any param") { + val org = mock[Complex] + + val counter = new AtomicInteger(1) + org.bar answers counter.getAndIncrement().toString + + counter.get should equal(1) + org.bar should be("1") + counter.get should equal(2) + org.bar should be("2") + } + + test("create a mock while stubbing another") { + val aMock = mock[Complex] + + aMock.returnBar returns mock[Bar] + aMock.returnBar should be(a[Bar]) + } + + test("stub a runtime exception instance to be thrown") { + val aMock = mock[Foo] + + aMock.bar(any[String]) throws new IllegalArgumentException + an[IllegalArgumentException] should be thrownBy aMock.bar("alice") + } + + test("stub a checked exception instance to be thrown") { + val aMock = mock[Foo] + + aMock.bar(any[String]) throws new Exception + an[Exception] should be thrownBy aMock.bar("fred") + } + + test("stub a multiple exceptions instance to be thrown") { + val aMock = mock[Complex] + + aMock.bar throws new Exception andThenThrow new IllegalArgumentException + + an[Exception] should be thrownBy aMock.bar + an[IllegalArgumentException] should be thrownBy aMock.bar + + aMock.returnBar throws new IllegalArgumentException andThenThrow new Exception + + an[IllegalArgumentException] should be thrownBy aMock.returnBar + an[Exception] should be thrownBy aMock.returnBar + } + + test("default answer should deal with default arguments") { + val aMock = mock[Complex] + + aMock.iHaveSomeDefaultArguments("I will not pass the second argument") + aMock.iHaveSomeDefaultArguments("I will pass the second argument", "second argument") + + MockitoSugar + .verify(aMock) + .iHaveSomeDefaultArguments("I will not pass the second argument", "default value") + MockitoSugar + .verify(aMock).iHaveSomeDefaultArguments("I will pass the second argument", "second argument") + } + + test("work with by-name arguments (argument order doesn't matter when not using matchers)") { + val aMock = mock[Complex] + + aMock.iStartWithByNameArgs("arg1", "arg2") returns "mocked!" + + aMock.iStartWithByNameArgs("arg1", "arg2") should be("mocked!") + aMock.iStartWithByNameArgs("arg111", "arg2") should not be "mocked!" + + MockitoSugar.verify(aMock).iStartWithByNameArgs("arg1", "arg2") + MockitoSugar.verify(aMock).iStartWithByNameArgs("arg111", "arg2") + } + + test("work with primitive by-name arguments") { + val aMock = mock[Complex] + + aMock.iHavePrimitiveByNameArgs(1, "arg2") returns "mocked!" + + aMock.iHavePrimitiveByNameArgs(1, "arg2") should be("mocked!") + aMock.iHavePrimitiveByNameArgs(2, "arg2") should not be "mocked!" + + MockitoSugar.verify(aMock).iHavePrimitiveByNameArgs(1, "arg2") + MockitoSugar.verify(aMock).iHavePrimitiveByNameArgs(2, "arg2") + } + + test("function0 arguments") { + val aMock = mock[Complex] + + aMock.iHaveFunction0Args(eqTo("arg1"), function0("arg2")) returns "mocked!" + + aMock.iHaveFunction0Args("arg1", () => "arg2") should be("mocked!") + aMock.iHaveFunction0Args("arg1", () => "arg3") should not be "mocked!" + + MockitoSugar.verify(aMock).iHaveFunction0Args(eqTo("arg1"), function0("arg2")) + MockitoSugar.verify(aMock).iHaveFunction0Args(eqTo("arg1"), function0("arg3")) + } + + test("ignore the calls to the methods that provide default arguments") { + val aMock = mock[Complex] + + aMock.iHaveSomeDefaultArguments("I will not pass the second argument") + + MockitoSugar + .verify(aMock) + .iHaveSomeDefaultArguments("I will not pass the second argument", "default value") + verifyNoMoreInteractions(aMock) + } + + test("support an ArgCaptor and deal with default arguments") { + val aMock = mock[Complex] + + aMock.iHaveSomeDefaultArguments("I will not pass the second argument") + + val captor1 = ArgCaptor[String] + val captor2 = ArgCaptor[String] + MockitoSugar.verify(aMock).iHaveSomeDefaultArguments(captor1, captor2) + + captor1 hasCaptured "I will not pass the second argument" + captor2 hasCaptured "default value" + } + + test("work with Future types 1") { + val mockFn = mock[() => Int] + mockFn() returns 42 + + Future(mockFn.apply()).map { v => + v should equal(42) + mockFn() was called + } + } + + test("work with Future types 2") { + val mockFutureFn = mock[() => Future[Int]] + mockFutureFn() returns Future(42) + + Await.result(mockFutureFn(), 2.seconds) should equal(42) + } + + test("works with Future whenReady 1") { + val mockFn = mock[() => Boolean] + mockFn() returns true + + val f: Future[Boolean] = Future.value(mockFn()) + whenReady(f) { _: Boolean => + mockFn() was called + } + } + + test("works with Future whenReady 2") { + val mockFutureFn = mock[() => Future[Int]] + mockFutureFn() returns Future(42) + + whenReady(mockFutureFn())(result => result should equal(42)) + } + + test("works with Future whenReady 3") { + val mockFn = mock[Int => Boolean] + mockFn(any[Int]) returns true + + val f: Future[Boolean] = Future.apply(mockFn(42)) + whenReady(f) { result: Boolean => + result should be(true) + } + } + + test("works with Function types") { + val expected = Future("Hello, world") + val mockFunction = mock[Int => Future[String]] + mockFunction(any[Int]) returns expected + + mockFunction(42) should equal(expected) + } + + test("works with Function types no sugar") { + val expected = Future("Hello, world") + val mockFunction = mock[Function1[Int, Future[String]]] + mockFunction(any[Int]) returns expected + + mockFunction(42) should equal(expected) + } + + test("works with Function whenReady") { + val mockFunction = mock[Int => Future[String]] + mockFunction(any[Int]) returns Future("Hello, world") + + whenReady(mockFunction(42))(result => result should equal("Hello, world")) + } + + test("works with Function whenReady no sugar") { + val mockFunction = mock[Function1[Int, Future[String]]] + mockFunction(any[Int]) returns Future("Hello, world") + + whenReady(mockFunction(42))(result => result should equal("Hello, world")) + } + + test("disallow passing traits in the settings") { + a[IllegalArgumentException] should be thrownBy { + mock[Foo](withSettings.extraInterfaces(classOf[Baz])) + } + } + + private final implicit val patienceConfig: PatienceConfig = + PatienceConfig(timeout = Span(5, Seconds), interval = Span(20, Millis)) + + private final implicit def convertTwitterFuture[T](twitterFuture: Future[T]): FutureConcept[T] = + new FutureConcept[T] { + def eitherValue: Option[Either[Throwable, T]] = + twitterFuture.poll.map { + case Return(result) => Right(result) + case Throw(e) => Left(e) + } + // As per the ScalaTest docs if the underlying future does not support these they + // must return false. + def isExpired: Boolean = false + def isCanceled: Boolean = false + } + + private final implicit def twitterDurationToScalaTestTimeout(duration: Duration): Timeout = { + Timeout(Span(duration.inMillis, Millis)) + } + + private final implicit def twitterDurationToScalaTestInterval(duration: Duration): Interval = { + Interval(Span(duration.inMillis, Millis)) + } +} diff --git a/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitWithFixtureTest.scala b/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitWithFixtureTest.scala new file mode 100644 index 0000000000..fdaaa2fee9 --- /dev/null +++ b/util-mock/src/test/scala/com/twitter/util/mock/MockitoTraitWithFixtureTest.scala @@ -0,0 +1,50 @@ +package com.twitter.util.mock + +import org.junit.runner.RunWith +import org.mockito.exceptions.verification.{NeverWantedButInvoked, NoInteractionsWanted} +import org.scalatest.Outcome +import org.scalatest.funsuite.FixtureAnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class MockitoTraitWithFixtureTest extends FixtureAnyFunSuite with Matchers with Mockito { + class Foo { + def bar(a: String) = "bar" + def baz = "baz" + } + + class FixtureParam { + val foo: Foo = mock[Foo] + } + + protected def withFixture(test: OneArgTest): Outcome = { + val theFixture = new FixtureParam + super.withFixture(test.toNoArgTest(theFixture)) + } + + test("verify no calls on fixture objects methods") { f: FixtureParam => + "mocked" willBe returned by f.foo.bar("jane") + "mocked" willBe returned by f.foo.baz + + f.foo.bar("jane") should be("mocked") + f.foo.baz should be("mocked") + + a[NeverWantedButInvoked] should be thrownBy { + f.foo.bar(any[String]) wasNever called + } + a[NeverWantedButInvoked] should be thrownBy { + f.foo.baz wasNever called + } + } + + test("verify no calls on a mock inside a fixture object") { f: FixtureParam => + f.foo.bar("jane") returns "mocked" + f.foo wasNever called + + f.foo.bar("jane") should be("mocked") + a[NoInteractionsWanted] should be thrownBy { + f.foo wasNever called + } + } +} diff --git a/util-mock/src/test/scala/com/twitter/util/mock/ResetMocksAfterEachTestTraitTest.scala b/util-mock/src/test/scala/com/twitter/util/mock/ResetMocksAfterEachTestTraitTest.scala new file mode 100644 index 0000000000..d62851065e --- /dev/null +++ b/util-mock/src/test/scala/com/twitter/util/mock/ResetMocksAfterEachTestTraitTest.scala @@ -0,0 +1,89 @@ +package com.twitter.util.mock + +import org.junit.runner.RunWith +import org.mockito.MockitoSugar +import org.mockito.scalatest.ResetMocksAfterEachTest +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatest.{BeforeAndAfterAll, BeforeAndAfterEach} +import org.scalatestplus.junit.JUnitRunner + +object ResetMocksAfterEachTestTraitTest { + + trait ReturnInteger { + def get: Int + } + + trait Foo { + def bar(a: String) = "bar" + } + + trait Baz { + def qux(a: String) = "qux" + } +} + +/** + * Ensure [[ResetMocksAfterEachTest]] works with [[com.twitter.util.mock.Mockito]] + */ +@RunWith(classOf[JUnitRunner]) +class ResetMocksAfterEachTestTraitTest + extends AnyFunSuite + with Matchers + with BeforeAndAfterAll + with BeforeAndAfterEach + with Mockito + with ResetMocksAfterEachTest { + import ResetMocksAfterEachTestTraitTest._ + + val foo: Foo = mock[Foo] + val baz: Baz = mock[Baz] + + private val mockReturnInteger: ReturnInteger = mock[ReturnInteger] + mockReturnInteger.get returns 100 + + override def beforeAll(): Unit = { + mockReturnInteger.get should equal(100) + super.beforeAll() + } + + override def afterEach(): Unit = { + super.afterEach() + + mockingDetails(mockReturnInteger).getInvocations.isEmpty should be(true) + } + + test("no interactions 1") { + MockitoSugar.verifyZeroInteractions(foo) + + foo.bar("bar") returns "mocked" + foo.bar("bar") shouldBe "mocked" + } + + test("no interactions 2") { + MockitoSugar.verifyZeroInteractions(foo) + + foo.bar("bar") returns "mocked2" + foo.bar("bar") shouldBe "mocked2" + } + + test("no interactions for all 1") { + MockitoSugar.verifyZeroInteractions(foo, baz) + + foo.bar("bar") returns "mocked3" + baz.qux("qux") returns "mocked4" + + foo.bar("bar") shouldBe "mocked3" + baz.qux("qux") shouldBe "mocked4" + } + + test("no interactions for all 2") { + MockitoSugar.verifyZeroInteractions(foo, baz) + + foo.bar("bar") returns "mocked5" + baz.qux("qux") returns "mocked6" + + foo.bar("bar") shouldBe "mocked5" + baz.qux("qux") shouldBe "mocked6" + } +} diff --git a/util-reflect/BUILD b/util-reflect/BUILD index e5ef1bc6f2..42fc74881d 100644 --- a/util-reflect/BUILD +++ b/util-reflect/BUILD @@ -1,16 +1,13 @@ target( - tags = { - # Legacy code that shouldn't be used for new services. - "deprecated", - }, + tags = ["bazel-compatible"], dependencies = [ "util/util-reflect/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ - "util/util-reflect/src/test/scala", ], ) diff --git a/util-reflect/src/main/scala/BUILD b/util-reflect/src/main/scala/BUILD index 31c0235a79..3282f48b4b 100644 --- a/util-reflect/src/main/scala/BUILD +++ b/util-reflect/src/main/scala/BUILD @@ -1,16 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-reflect", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/cglib", - "util/util-core/src/main/scala", - ], - exports = [ - "3rdparty/jvm/cglib", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", ], ) diff --git a/util-reflect/src/main/scala/com/twitter/util/reflect/Annotations.scala b/util-reflect/src/main/scala/com/twitter/util/reflect/Annotations.scala new file mode 100644 index 0000000000..69467c4559 --- /dev/null +++ b/util-reflect/src/main/scala/com/twitter/util/reflect/Annotations.scala @@ -0,0 +1,290 @@ +package com.twitter.util.reflect + +import com.twitter.util.{Return, Throw, Try} +import java.lang.annotation.Annotation +import java.lang.reflect.{Field, Parameter} +import scala.reflect.{ClassTag, classTag} + +/** Utility methods for dealing with [[java.lang.annotation.Annotation]] */ +object Annotations { + + /** + * Filter a list of annotations discriminated on whether the + * annotation is itself annotated with the passed annotation. + * + * @param annotations the list of [[Annotation]] instances to filter. + * @tparam A the type of the annotation by which to filter. + * + * @return the filtered list of matching annotations. + */ + def filterIfAnnotationPresent[A <: Annotation: ClassTag]( + annotations: Array[Annotation] + ): Array[Annotation] = annotations.filter(isAnnotationPresent[A]) + + /** + * Filters a list of annotations by annotation type discriminated by the set of given annotations. + * @param filterSet the Set of [[Annotation]] classes by which to filter. + * @param annotations the list [[Annotation]] instances to filter. + * + * @return the filtered list of matching annotations. + */ + def filterAnnotations( + filterSet: Set[Class[_ <: Annotation]], + annotations: Array[Annotation] + ): Array[Annotation] = + annotations.filter(a => filterSet.contains(a.annotationType)) + + /** + * Find an [[Annotation]] within a given list of annotations of the given target. + * annotation type. + * @param target the class of the [[Annotation]] to find. + * @param annotations the list of [[Annotation]] instances to search. + * + * @return the matching [[Annotation]] instance if found, otherwise None. + */ + // optimized + def findAnnotation( + target: Class[_ <: Annotation], + annotations: Array[Annotation] + ): Option[Annotation] = { + var found: Option[Annotation] = None + var index = 0 + while (index < annotations.length && found.isEmpty) { + val annotation = annotations(index) + if (annotation.annotationType() == target) found = Some(annotation) + index += 1 + } + found + } + + /** + * Find an [[Annotation]] within a given list of annotations annotated by the given type param. + * @param annotations the list of [[Annotation]] instances to search. + * @tparam A the type of the [[Annotation]] to find. + * + * @return the matching [[Annotation]] instance if found, otherwise None. + */ + // optimized + def findAnnotation[A <: Annotation: ClassTag](annotations: Array[Annotation]): Option[A] = { + val size = annotations.length + val annotationType = classTag[A].runtimeClass.asInstanceOf[Class[A]] + var found: Option[A] = None + var index = 0 + while (found.isEmpty && index < size) { + val annotation = annotations(index) + if (annotation.annotationType() == annotationType) found = Some(annotation.asInstanceOf[A]) + index += 1 + } + found + } + + /** + * Attempts to map fields to array of annotations. This is intended to work around the Scala default + * of not carrying annotation information on annotated constructor fields annotated without also + * specifying a @meta annotation (e.g., `@field`). Thus, this method scans the constructor and + * then all declared fields in order to build up the mapping of field name to `Array[java.lang.Annotation]`. + * When the given class has a single constructor, that constructor will be used. Otherwise, the + * given parameter types are used to locate a constructor. + * Steps: + * - walk the constructor and collect annotations on all parameters. + * - for each declared field in the class, collect declared annotations. If a constructor parameter + * is overridden as a declared field, the annotations on the declared field take precedence. + * - for each super interface, walk each declared method, collect declared method annotations IF + * the declared method is the same name of a declared field from step two. That is find only + * super interface methods which are implemented by declared fields in the given class in order + * to locate inherited annotations for the declared field. + * + * @param clazz the `Class` to inspect. This should represent a Scala case class. + * @param parameterTypes an optional list of parameter class types in order to locate the appropriate + * case class constructor when the class has multiple constructors. If a class + * has multiple constructors and parameter types are not specified, an + * [[IllegalArgumentException]] is thrown. Likewise, if a suitable constructor + * cannot be located by the given parameter types an [[IllegalArgumentException]] + * is thrown. + * + * @see [[https://www.scala-lang.org/api/current/scala/annotation/meta/index.html]] + * @return a map of field name to associated annotations. An entry in the map only occurs if the + * associated annotations are non-empty. That is fields without any found annotations are + * not returned. + */ + // optimized + def findAnnotations( + clazz: Class[_], + parameterTypes: Seq[Class[_]] = Seq.empty + ): scala.collection.Map[String, Array[Annotation]] = { + val clazzAnnotations = scala.collection.mutable.HashMap[String, Array[Annotation]]() + val constructorCount = clazz.getConstructors.length + val constructor = if (constructorCount > 1) { + if (parameterTypes.isEmpty) + throw new IllegalArgumentException( + "Multiple constructors detected and no parameter types specified.") + Try(clazz.getConstructor(parameterTypes: _*)) match { + case Return(con) => con + case Throw(t: NoSuchMethodException) => + // translate from the reflection exception to a runtime exception + throw new IllegalArgumentException(t) + case Throw(t) => throw t + } + } else clazz.getConstructors.head + + val fields: Array[Parameter] = constructor.getParameters + val clazzConstructorAnnotations: Array[Array[Annotation]] = constructor.getParameterAnnotations + + var i = 0 + while (i < fields.length) { + val field = fields(i) + val fieldAnnotations = clazzConstructorAnnotations(i) + if (fieldAnnotations.nonEmpty) clazzAnnotations.put(field.getName, fieldAnnotations) + i += 1 + } + + val declaredFields: Array[Field] = clazz.getDeclaredFields + var j = 0 + while (j < declaredFields.length) { + val field = declaredFields(j) + if (field.getDeclaredAnnotations.nonEmpty) { + if (!clazzAnnotations.contains(field.getName)) { + clazzAnnotations.put(field.getName, field.getDeclaredAnnotations) + } else { + val existing = clazzAnnotations(field.getName) + // prefer field annotations over constructor annotation + mergeAnnotationArrays(field.getDeclaredAnnotations, existing) + } + } + + j += 1 + } + + // find inherited annotations for declared fields + findDeclaredAnnotations( + clazz, + declaredFields.map(_.getName).toSet, + clazzAnnotations + ) + } + + /** + * Determines if the given [[Annotation]] has an annotation type of the given type param. + * @param annotation the [[Annotation]] to match. + * @tparam A the type to match against. + * + * @return true if the given [[Annotation]] is of type [[A]], false otherwise. + */ + def equals[A <: Annotation: ClassTag](annotation: Annotation): Boolean = + annotation.annotationType() == classTag[A].runtimeClass.asInstanceOf[Class[A]] + + /** + * Determines if the given [[Annotation]] is annotated by an [[Annotation]] of the given + * type param. + * @param annotation the [[Annotation]] to match. + * @tparam A the type of the [[Annotation]] to determine if is annotated on the [[Annotation]]. + * + * @return true if the given [[Annotation]] is annotated with an [[Annotation]] of type [[A]], + * false otherwise. + */ + def isAnnotationPresent[A <: Annotation: ClassTag](annotation: Annotation): Boolean = + annotation.annotationType.isAnnotationPresent(classTag[A].runtimeClass.asInstanceOf[Class[A]]) + + /** + * Determines if the given [[A]] is annotated by an [[Annotation]] of the given + * type param [[ToFindAnnotation]]. + * + * @tparam ToFindAnnotation the [[Annotation]] to match. + * @tparam A the type of the [[Annotation]] to determine if is annotated on the [[A]]. + * + * @return true if the given [[Annotation]] is annotated with an [[Annotation]] of type [[ToFindAnnotation]], + * false otherwise. + */ + def isAnnotationPresent[ + ToFindAnnotation <: Annotation: ClassTag, + A <: Annotation: ClassTag + ]: Boolean = { + val annotationToFindClazz: Class[Annotation] = + classTag[ToFindAnnotation].runtimeClass + .asInstanceOf[Class[Annotation]] + + classTag[A].runtimeClass + .asInstanceOf[Class[A]].isAnnotationPresent(annotationToFindClazz) + } + + /** + * Attempts to return the result of invoking `method()` of a given [[Annotation]] which should be + * annotated by an [[Annotation]] of type [[A]]. + * + * If the given [[Annotation]] is not annotated by an [[Annotation]] of type [[A]] or if the + * `Annotation#method` cannot be invoked then [[None]] is returned. + * + * @param annotation the [[Annotation]] to process. + * @param method the method of the [[Annotation]] to invoke to obtain a value. + * @tparam A the [[Annotation]] type to match against before trying to invoke `annotation#method`. + * + * @return the result of invoking `annotation#method()` or None. + */ + def getValueIfAnnotatedWith[A <: Annotation: ClassTag]( + annotation: Annotation, + method: String = "value" + ): Option[String] = { + if (isAnnotationPresent[A](annotation)) { + getValue(annotation, method) + } else None + } + + /** + * Attempts to return the result of invoking `method()` of a given [[Annotation]]. If the + * `Annotation#method` cannot be invoked then [[None]] is returned. + * + * @param annotation the [[Annotation]] to process. + * @param method the method of the [[Annotation]] to invoke to obtain a value. + * + * @return the result of invoking `annotation#method()` or None. + */ + def getValue( + annotation: Annotation, + method: String = "value" + ): Option[String] = { + for { + method <- annotation.getClass.getDeclaredMethods.find(_.getName == method) + value <- Try(method.invoke(annotation).asInstanceOf[String]).toOption + } yield value + } + + // optimized + private[twitter] def findDeclaredAnnotations( + clazz: Class[_], + declaredFields: Set[String], // `.contains` on a Set should be faster than on an Array + acc: scala.collection.mutable.Map[String, Array[Annotation]] + ): scala.collection.Map[String, Array[Annotation]] = { + val methods = clazz.getDeclaredMethods + var i = 0 + while (i < methods.length) { + val method = methods(i) + val methodAnnotations = method.getDeclaredAnnotations + if (methodAnnotations.nonEmpty && declaredFields.contains(method.getName)) { + acc.get(method.getName) match { + case Some(existing) => + acc.put(method.getName, mergeAnnotationArrays(existing, methodAnnotations)) + case _ => + acc.put(method.getName, methodAnnotations) + } + } + i += 1 + } + + val interfaces = clazz.getInterfaces + var j = 0 + while (j < interfaces.length) { + val interface = interfaces(j) + findDeclaredAnnotations(interface, declaredFields, acc) + j += 1 + } + + acc + } + + /** Prefer values in A over B */ + private[this] def mergeAnnotationArrays( + a: Array[Annotation], + b: Array[Annotation] + ): Array[Annotation] = + a ++ b.filterNot(bAnnotation => a.exists(_.annotationType() == bAnnotation.annotationType())) +} diff --git a/util-reflect/src/main/scala/com/twitter/util/reflect/BUILD b/util-reflect/src/main/scala/com/twitter/util/reflect/BUILD new file mode 100644 index 0000000000..5b38a2e602 --- /dev/null +++ b/util-reflect/src/main/scala/com/twitter/util/reflect/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-reflect", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-reflect/src/main/scala/com/twitter/util/reflect/Classes.scala b/util-reflect/src/main/scala/com/twitter/util/reflect/Classes.scala new file mode 100644 index 0000000000..6edbf9bb26 --- /dev/null +++ b/util-reflect/src/main/scala/com/twitter/util/reflect/Classes.scala @@ -0,0 +1,26 @@ +package com.twitter.util.reflect + +import com.twitter.util.Memoize + +object Classes { + private val PRODUCT: Class[Product] = classOf[Product] + private val OPTION: Class[Option[_]] = classOf[Option[_]] + private val LIST: Class[List[_]] = classOf[List[_]] + + /** + * Safely compute the `clazz.getSimpleName` of a given [[Class]] handling malformed + * names with a best-effort to parse the final segment after the last `.` after + * transforming all `$` to `.`. + */ + val simpleName: Class[_] => String = Memoize { clazz: Class[_] => + // class.getSimpleName can fail with an java.lan.InternalError (for malformed class/package names), + // so we attempt to deal with this with a manual translation if necessary + try { + clazz.getSimpleName + } catch { + case _: InternalError => + // replace all $ with . and return the last element + clazz.getName.replace('$', '.').split('.').last + } + } +} diff --git a/util-reflect/src/main/scala/com/twitter/util/reflect/Proxy.scala b/util-reflect/src/main/scala/com/twitter/util/reflect/Proxy.scala deleted file mode 100644 index 5b9301e099..0000000000 --- a/util-reflect/src/main/scala/com/twitter/util/reflect/Proxy.scala +++ /dev/null @@ -1,121 +0,0 @@ -package com.twitter.util.reflect - -import java.io.Serializable -import java.lang.reflect.Method - -import net.sf.cglib.proxy.{MethodInterceptor => CGMethodInterceptor, _} - -import com.twitter.util.Future - -class NonexistentTargetException extends Exception("MethodCall was invoked without a valid target.") - -object Proxy { - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply[I <: AnyRef: Manifest](f: MethodCall[I] => AnyRef) = { - new ProxyFactory[I](f).apply() - } - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply[I <: AnyRef: Manifest](target: I, f: MethodCall[I] => AnyRef) = { - new ProxyFactory[I](f).apply(target) - } -} - -object ProxyFactory { - private[reflect] object NoOpInterceptor extends MethodInterceptor[AnyRef](None, m => null) - - private[reflect] object IgnoredMethodFilter extends CallbackFilter { - def accept(m: Method) = { - m.getName match { - case "hashCode" => 1 - case "equals" => 1 - case "toString" => 1 - case _ => 0 - } - } - } -} - -class AbstractProxyFactory[I <: AnyRef: Manifest] { - import ProxyFactory._ - - final val interface = implicitly[Manifest[I]].runtimeClass - - protected final val proto = { - val e = new Enhancer - e.setCallbackFilter(IgnoredMethodFilter) - e.setCallbacks(Array(NoOpInterceptor, NoOp.INSTANCE)) - e.setSuperclass(interface) - e.create.asInstanceOf[Factory] - } - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - protected final def newWithCallback(f: MethodCall[I] => AnyRef) = { - proto.newInstance(Array(new MethodInterceptor(None, f), NoOp.INSTANCE)).asInstanceOf[I] - } - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - protected final def newWithCallback[T <: I](target: T, f: MethodCall[I] => AnyRef) = { - proto.newInstance(Array(new MethodInterceptor(Some(target), f), NoOp.INSTANCE)).asInstanceOf[I] - } -} - -class ProxyFactory[I <: AnyRef: Manifest](f: MethodCall[I] => AnyRef) - extends AbstractProxyFactory[I] { - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply[T <: I](target: T) = newWithCallback(target, f) - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply() = newWithCallback(f) -} - -private[reflect] class MethodInterceptor[I <: AnyRef]( - target: Option[I], - callback: MethodCall[I] => AnyRef -) extends CGMethodInterceptor - with Serializable { - val targetRef = target.getOrElse(null).asInstanceOf[I] - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - final def intercept(p: AnyRef, m: Method, args: Array[AnyRef], methodProxy: MethodProxy) = { - callback(new MethodCall(targetRef, m, args, methodProxy)) - } -} - -final class MethodCall[T <: AnyRef] private[reflect] ( - targetRef: T, - val method: Method, - val args: Array[AnyRef], - methodProxy: MethodProxy -) extends (() => AnyRef) { - - lazy val target = if (targetRef ne null) Some(targetRef) else None - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def clazz = method.getDeclaringClass - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def clazzName = clazz.getName - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def className = clazzName - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def parameterTypes = method.getParameterTypes - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def name = method.getName - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def returnsUnit = { - val rt = method.getReturnType - (rt eq classOf[Unit]) || (rt eq classOf[Null]) || (rt eq java.lang.Void.TYPE) - } - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def returnsFuture = classOf[Future[_]] isAssignableFrom method.getReturnType - - private def getTarget = if (targetRef ne null) targetRef else throw new NonexistentTargetException - - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply() = methodProxy.invoke(getTarget, args) - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply(newTarget: T) = methodProxy.invoke(newTarget, args) - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply(newArgs: Array[AnyRef]) = methodProxy.invoke(getTarget, newArgs) - @deprecated("Legacy code that shouldn't be used for new services", "2018-05-25") - def apply(newTarget: T, newArgs: Array[AnyRef]) = methodProxy.invoke(newTarget, newArgs) -} diff --git a/util-reflect/src/main/scala/com/twitter/util/reflect/Types.scala b/util-reflect/src/main/scala/com/twitter/util/reflect/Types.scala new file mode 100644 index 0000000000..0f0b1662b7 --- /dev/null +++ b/util-reflect/src/main/scala/com/twitter/util/reflect/Types.scala @@ -0,0 +1,126 @@ +package com.twitter.util.reflect + +import com.twitter.util.Memoize +import java.lang.reflect.{ParameterizedType, TypeVariable, Type => JavaType} +import scala.reflect.api.TypeCreator +import scala.reflect.runtime.universe._ + +object Types { + private[this] val PRODUCT: Type = typeOf[Product] + private[this] val OPTION: Type = typeOf[Option[_]] + private[this] val LIST: Type = typeOf[List[_]] + + /** + * Returns `true` if the given class type is considered a case class. + * True if: + * - is assignable from PRODUCT and + * - not assignable from OPTION nor LIST and + * - is not a Tuple and + * - class symbol is case class. + */ + val isCaseClass: Class[_] => Boolean = Memoize { clazz: Class[_] => + val tpe = asTypeTag(clazz).tpe + val classSymbol = tpe.typeSymbol.asClass + tpe <:< PRODUCT && + !(tpe <:< OPTION || tpe <:< LIST) && + !clazz.getName.startsWith("scala.Tuple") && + classSymbol.isCaseClass + } + + /** + * This is the negation of [[Types.isCaseClass]] + * Determine if a given class type is not a case class. + * Returns `true` if it is NOT considered a case class. + */ + def notCaseClass[T](clazz: Class[T]): Boolean = !isCaseClass(clazz) + + /** + * Convert from the given `Class[T]` to a `TypeTag[T]` in the runtime universe. + * =Usage= + * {{{ + * val clazz: Class[T] = ??? + * val tag: TypeTag[T] = Types.asTypeTag(clazz) + * }}} + * + * @param clazz the class for which to build the resultant [[TypeTag]] + * @return a `TypeTag[T]` representing the given `Class[T]`. + */ + def asTypeTag[T](clazz: Class[_ <: T]): TypeTag[T] = { + val clazzMirror = runtimeMirror(clazz.getClassLoader) + val tpe = clazzMirror.classSymbol(clazz).toType + val typeCreator = new TypeCreator() { + def apply[U <: scala.reflect.api.Universe with scala.Singleton]( + m: scala.reflect.api.Mirror[U] + ): U#Type = { + if (clazzMirror != m) throw new RuntimeException("wrong mirror") + else tpe.asInstanceOf[U#Type] + } + } + TypeTag[T](clazzMirror, typeCreator) + } + + /** + * Return the runtime class from a given [[TypeTag]]. + * =Usage= + * {{{ + * val clazz: Class[T] = Types.runtimeClass[T] + * }}} + * + * or pass in the implict TypeTag explicitly: + * + * {{{ + * val tag: TypeTag[T] = ??? + * val clazz: Class[T] = Types.runtimeClass(tag) + * }}} + * + * + * @note the given [[TypeTag]] must be from the runtime universe otherwise an + * [[IllegalArgumentException]] will be thrown. + * + * @tparam T the [[TypeTag]] and expected class type of the returned class. + * @return the runtime class of the given [[TypeTag]]. + */ + def runtimeClass[T](implicit tag: TypeTag[T]): Class[T] = { + val clazzMirror = runtimeMirror(getClass.getClassLoader) + if (clazzMirror != tag.mirror) { + throw new IllegalArgumentException("TypeTag is not from runtime universe.") + } + tag.mirror.runtimeClass(tag.tpe).asInstanceOf[Class[T]] + } + + /** + * If `thisTypeTag` [[TypeTag]] has the same TypeSymbol as `thatTypeTag` [[TypeTag]]. + * + * @param thatTypeTag [[TypeTag]] to compare + * @param thisTypeTag [[TypeTag]] to compare + * @tparam T type of thisTypeTag + * @return true if the TypeSymbols are equivalent. + * @see [[scala.reflect.api.Symbols]] + */ + def equals[T](thatTypeTag: TypeTag[_])(implicit thisTypeTag: TypeTag[T]): Boolean = + thisTypeTag.tpe.typeSymbol == thatTypeTag.tpe.typeSymbol + + /** + * If the given [[java.lang.reflect.Type]] is parameterized, return an Array of the + * type parameter names. E.g., `Map[T, U]` returns, `Array("T", "U")`. + * + * The use case is when we are trying to match the given type's type parameters to + * a set of bindings that are stored keyed by the type name, e.g. if the given type + * is a `Map[K, V]` we want to be able to look up the binding for the key `K` at runtime + * during reflection operations e.g., if the K type is bound to a String we want to + * be able use that when further processing this type. This type of operation most typically + * happens with Jackson reflection handling of case classes, which is a specialized case + * and thus the utility of this method may not be broadly applicable and therefore is limited + * in visibility. + */ + private[twitter] def parameterizedTypeNames(javaType: JavaType): Array[String] = + javaType match { + case parameterizedType: ParameterizedType => + parameterizedType.getActualTypeArguments.map(_.getTypeName) + case typeVariable: TypeVariable[_] => + Array(typeVariable.getTypeName) + case clazz: Class[_] => + clazz.getTypeParameters.map(_.getName) + case _ => Array.empty + } +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Annotation1.java b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation1.java new file mode 100644 index 0000000000..19be3a95f6 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation1.java @@ -0,0 +1,11 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target({ElementType.ANNOTATION_TYPE, ElementType.METHOD, ElementType.FIELD, ElementType.PARAMETER}) +@Retention(RetentionPolicy.RUNTIME) +public @interface Annotation1 { +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Annotation2.java b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation2.java new file mode 100644 index 0000000000..5744fbf3c5 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation2.java @@ -0,0 +1,12 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target({ElementType.ANNOTATION_TYPE, ElementType.METHOD, ElementType.FIELD, ElementType.PARAMETER}) +@Retention(RetentionPolicy.RUNTIME) +@MarkerAnnotation +public @interface Annotation2 { +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Annotation3.java b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation3.java new file mode 100644 index 0000000000..8ab9103703 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation3.java @@ -0,0 +1,14 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target({ElementType.ANNOTATION_TYPE, ElementType.METHOD, ElementType.FIELD, ElementType.PARAMETER}) +@Retention(RetentionPolicy.RUNTIME) +@MarkerAnnotation +public @interface Annotation3 { + /** A value */ + String value() default "annotation3"; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Annotation4.java b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation4.java new file mode 100644 index 0000000000..832c335f56 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation4.java @@ -0,0 +1,11 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target({ElementType.ANNOTATION_TYPE, ElementType.METHOD, ElementType.FIELD, ElementType.PARAMETER}) +@Retention(RetentionPolicy.RUNTIME) +public @interface Annotation4 { +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Annotation5.java b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation5.java new file mode 100644 index 0000000000..2d378948d8 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Annotation5.java @@ -0,0 +1,13 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target({ElementType.ANNOTATION_TYPE, ElementType.METHOD, ElementType.FIELD, ElementType.PARAMETER}) +@Retention(RetentionPolicy.RUNTIME) +public @interface Annotation5 { + /** A value */ + String discriminator() default "annotation5"; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/BUILD.bazel b/util-reflect/src/test/java/com/twitter/util/reflect/BUILD.bazel new file mode 100644 index 0000000000..fb279e4d23 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/BUILD.bazel @@ -0,0 +1,7 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], +) diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/MarkerAnnotation.java b/util-reflect/src/test/java/com/twitter/util/reflect/MarkerAnnotation.java new file mode 100644 index 0000000000..cde33f2bbe --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/MarkerAnnotation.java @@ -0,0 +1,10 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Target(ElementType.TYPE) +@Retention(RetentionPolicy.RUNTIME) +public @interface MarkerAnnotation {} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Thing.java b/util-reflect/src/test/java/com/twitter/util/reflect/Thing.java new file mode 100644 index 0000000000..e9b05d495e --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Thing.java @@ -0,0 +1,15 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Retention(RUNTIME) +@Target({ElementType.FIELD, ElementType.PARAMETER, ElementType.METHOD}) +@MarkerAnnotation +public @interface Thing { + /** Name of the thing */ + String value() default "thing"; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/ThingImpl.java b/util-reflect/src/test/java/com/twitter/util/reflect/ThingImpl.java new file mode 100644 index 0000000000..f7d3260068 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/ThingImpl.java @@ -0,0 +1,43 @@ +package com.twitter.util.reflect; + +import java.io.Serializable; +import java.lang.annotation.Annotation; + +class ThingImpl implements Thing, Serializable { + private final String value; + + public ThingImpl(String value) { + this.value = value; + } + + public String value() { + return this.value; + } + + public int hashCode() { + // This is specified in java.lang.Annotation. + return (127 * "value".hashCode()) ^ value.hashCode(); + } + + /** + * Thing specific equals + */ + public boolean equals(Object o) { + if (!(o instanceof Thing)) { + return false; + } + + Thing other = (Thing) o; + return value.equals(other.value()); + } + + public String toString() { + return "@" + Thing.class.getName() + "(value=" + value + ")"; + } + + public Class annotationType() { + return Thing.class; + } + + private static final long serialVersionUID = 0; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Things.java b/util-reflect/src/test/java/com/twitter/util/reflect/Things.java new file mode 100644 index 0000000000..f9ec6f8e7b --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Things.java @@ -0,0 +1,11 @@ +package com.twitter.util.reflect; + +public final class Things { + private Things() { + } + + /** Creates a {@link Thing} annotation with {@code name} as the value. */ + public static Thing named(String name) { + return new ThingImpl(name); + } +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Widget.java b/util-reflect/src/test/java/com/twitter/util/reflect/Widget.java new file mode 100644 index 0000000000..d19080e651 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Widget.java @@ -0,0 +1,15 @@ +package com.twitter.util.reflect; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Retention(RUNTIME) +@Target({ElementType.FIELD, ElementType.PARAMETER, ElementType.METHOD}) +public @interface Widget { + + /** Name of the widget */ + String value() default "widget"; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/WidgetImpl.java b/util-reflect/src/test/java/com/twitter/util/reflect/WidgetImpl.java new file mode 100644 index 0000000000..284e5f809d --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/WidgetImpl.java @@ -0,0 +1,43 @@ +package com.twitter.util.reflect; + +import java.io.Serializable; +import java.lang.annotation.Annotation; + +class WidgetImpl implements Widget, Serializable { + private final String value; + + public WidgetImpl(String value) { + this.value = value; + } + + public String value() { + return this.value; + } + + public int hashCode() { + // This is specified in java.lang.Annotation. + return (127 * "value".hashCode()) ^ value.hashCode(); + } + + /** + * Widget specific equals + */ + public boolean equals(Object o) { + if (!(o instanceof Widget)) { + return false; + } + + Widget other = (Widget) o; + return value.equals(other.value()); + } + + public String toString() { + return "@" + Widget.class.getName() + "(value=" + value + ")"; + } + + public Class annotationType() { + return Widget.class; + } + + private static final long serialVersionUID = 0; +} diff --git a/util-reflect/src/test/java/com/twitter/util/reflect/Widgets.java b/util-reflect/src/test/java/com/twitter/util/reflect/Widgets.java new file mode 100644 index 0000000000..ca9ae572b6 --- /dev/null +++ b/util-reflect/src/test/java/com/twitter/util/reflect/Widgets.java @@ -0,0 +1,11 @@ +package com.twitter.util.reflect; + +public final class Widgets { + private Widgets() { + } + + /** Creates a {@link Widget} annotation with {@code name} as the value. */ + public static Widget named(String name) { + return new WidgetImpl(name); + } +} diff --git a/util-reflect/src/test/scala/BUILD b/util-reflect/src/test/scala/BUILD deleted file mode 100644 index b944a1382f..0000000000 --- a/util-reflect/src/test/scala/BUILD +++ /dev/null @@ -1,10 +0,0 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-reflect/src/main/scala", - ], -) diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/AnnotationsTest.scala b/util-reflect/src/test/scala/com/twitter/util/reflect/AnnotationsTest.scala new file mode 100644 index 0000000000..096c0d4f95 --- /dev/null +++ b/util-reflect/src/test/scala/com/twitter/util/reflect/AnnotationsTest.scala @@ -0,0 +1,390 @@ +package com.twitter.util.reflect + +import com.twitter.util.reflect.testclasses._ +import java.lang.annotation.Annotation +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import scala.collection.Map + +@RunWith(classOf[JUnitRunner]) +class AnnotationsTest extends AnyFunSuite with Matchers { + + test("Annotations#filterIfAnnotationPresent") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassOneTwo]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + + val found = Annotations.filterIfAnnotationPresent[MarkerAnnotation](annotations) + found.length should be(1) + found.head.annotationType should equal(classOf[Annotation2]) + } + + test("Annotations#filterAnnotations") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassThreeFour]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val filterSet: Set[Class[_ <: Annotation]] = Set(classOf[Annotation4]) + val found = Annotations.filterAnnotations(filterSet, annotations) + found.length should be(1) + found.head.annotationType should equal(classOf[Annotation4]) + } + + test("Annotations#findAnnotation") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassThreeFour]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation(classOf[Annotation1], annotations) should be(None) // not found + val found = Annotations.findAnnotation(classOf[Annotation3], annotations) + found.isDefined should be(true) + found.get.annotationType() should equal(classOf[Annotation3]) + } + + test("Annotations#findAnnotation by type") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassOneTwoThreeFour]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation2](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation3](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation4](annotations).isDefined should be(true) + } + + test("Annotations#equals") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassOneTwoThreeFour]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val found = Annotations.findAnnotation[Annotation1](annotations) + found.isDefined should be(true) + + val annotation = found.get + Annotations.equals[Annotation1](annotation) should be(true) + } + + test("Annotations#isAnnotationPresent") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassOneTwoThreeFour]) + annotationMap.isEmpty should be(false) + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val annotation1 = Annotations.findAnnotation[Annotation1](annotations).get + val annotation2 = Annotations.findAnnotation[Annotation2](annotations).get + val annotation3 = Annotations.findAnnotation[Annotation3](annotations).get + val annotation4 = Annotations.findAnnotation[Annotation4](annotations).get + + Annotations.isAnnotationPresent[MarkerAnnotation](annotation1) should be(false) + Annotations.isAnnotationPresent[MarkerAnnotation](annotation2) should be(true) + Annotations.isAnnotationPresent[MarkerAnnotation](annotation3) should be(true) + Annotations.isAnnotationPresent[MarkerAnnotation](annotation4) should be(false) + } + + test("Annotations#findAnnotations") { + var found: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[WithThings]) + found.isEmpty should be(false) + var annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(4) + + found = Annotations.findAnnotations(classOf[WithWidgets]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(4) + + found = Annotations.findAnnotations(classOf[CaseClassOneTwo]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(2) + + found = Annotations.findAnnotations(classOf[CaseClassThreeFour]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(2) + + found = Annotations.findAnnotations(classOf[CaseClassOneTwoThreeFour]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(4) + + found = Annotations.findAnnotations(classOf[CaseClassOneTwoWithFields]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(2) + + found = Annotations.findAnnotations(classOf[CaseClassOneTwoWithAnnotatedField]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(3) + + found = Annotations.findAnnotations(classOf[CaseClassThreeFourAncestorOneTwo]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(4) + + found = Annotations.findAnnotations(classOf[CaseClassAncestorOneTwo]) + found.isEmpty should be(false) + annotations = found.flatMap { case (_, annotations) => annotations.toList }.toArray + annotations.length should equal(3) + } + + test("Annotations.findAnnotations error") { + // parameter types are ignored + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations( + classOf[CaseClassOneTwoWithFields], + Seq(classOf[Boolean], classOf[Double])) + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + } + + test("Annotations#getValueIfAnnotatedWith") { + val annotationsMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[WithThings]) + annotationsMap.isEmpty should be(false) + val annotations = annotationsMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val annotation1 = Annotations.findAnnotation[Annotation1](annotations).get + val annotation2 = Annotations.findAnnotation[Annotation2](annotations).get + + // @Annotation1 is not annotated with @MarkerAnnotation + Annotations.getValueIfAnnotatedWith[MarkerAnnotation](annotation1).isDefined should be(false) + // @Annotation2 is annotated with @MarkerAnnotation but does not define a value() method + Annotations.getValueIfAnnotatedWith[MarkerAnnotation](annotation2).isDefined should be(false) + + val things = Annotations.filterAnnotations(Set(classOf[Thing]), annotations) + things.foreach { thing => + Annotations.getValueIfAnnotatedWith[MarkerAnnotation](thing).isDefined should be(true) + } + + val fooThingAnnotation = Things.named("foo") + val result = Annotations.getValueIfAnnotatedWith[MarkerAnnotation](fooThingAnnotation) + result.isDefined should be(true) + result.get should equal("foo") + } + + test("Annotations#getValue") { + val annotationsMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[CaseClassThreeFour]) + annotationsMap.isEmpty should be(false) + val annotations = annotationsMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val annotation3 = Annotations.findAnnotation[Annotation3](annotations).get + val annotation4 = Annotations.findAnnotation[Annotation4](annotations).get + + val annotation3Value = Annotations.getValue(annotation3) + annotation3Value.isDefined should be(true) + annotation3Value.get should equal("annotation3") + + // Annotation4 has no value + Annotations.getValue(annotation4).isDefined should be(false) + + val clazzFiveAnnotationsMap = Annotations.findAnnotations(classOf[CaseClassFive]) + val clazzFiveAnnotations = clazzFiveAnnotationsMap.flatMap { + case (_, annotations) => annotations.toList + }.toArray + val annotation5 = + Annotations.findAnnotation[Annotation5](clazzFiveAnnotations).get + // Annotation5 has a uniquely named method + val annotation5Value = Annotations.getValue(annotation5, "discriminator") + annotation5Value.isDefined should be(true) + annotation5Value.get should equal("annotation5") + + // directly look up a value on an annotation + val fooThingAnnotation = Things.named("foo") + val fooThingAnnotationValue = Annotations.getValue(fooThingAnnotation) + fooThingAnnotationValue.isDefined should be(true) + fooThingAnnotationValue.get should equal("foo") + } + + test("Annotations#secondaryConstructor") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations( + classOf[WithSecondaryConstructor], + Seq(classOf[String], classOf[String])) + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("three").length should equal(1) + annotationMap("three").head.annotationType() should equal(classOf[Annotation3]) + annotationMap("four").length should equal(1) + annotationMap("four").head.annotationType() should equal(classOf[Annotation4]) + annotationMap.get("one") should be(None) // not found + annotationMap.get("two") should be(None) // not found + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation2](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation3](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation4](annotations).isDefined should be(true) + } + + test("Annotations#secondaryConstructor 1") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations( + classOf[WithSecondaryConstructor], + Seq(classOf[Int], classOf[Int])) + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + annotationMap.get("three") should be(None) // not found + annotationMap.get("four") should be(None) // not found + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation2](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation3](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation4](annotations) should be(None) // not found + } + + test("Annotations#secondaryConstructor 2") { + val annotationMap: Map[String, Array[Annotation]] = Annotations.findAnnotations( + classOf[StaticSecondaryConstructor], + Seq(classOf[String], classOf[String]) + ) + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + } + + test("Annotations#secondaryConstructor 3") { + val annotationMap: Map[String, Array[Annotation]] = Annotations.findAnnotations( + classOf[StaticSecondaryConstructor], + Seq(classOf[Int], classOf[Int]) + ) + + // secondary int "constructor" is not used/scanned as it is an Object apply method which + // is a factory method and not a constructor. + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + } + + test("Annotations#secondaryConstructor 4") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[StaticSecondaryConstructor]) + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation2](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation3](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation4](annotations) should be(None) // not found + } + + test("Annotations#secondaryConstructor 5") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[StaticSecondaryConstructorWithMethodAnnotation]) + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + annotationMap.get("widget1") should be(None) + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + val annotation1 = Annotations.findAnnotation[Annotation1](annotations).get + val annotation2 = Annotations.findAnnotation[Annotation2](annotations).get + + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations) should be(Some(annotation1)) + Annotations.findAnnotation[Annotation2](annotations) should be(Some(annotation2)) + Annotations.findAnnotation[Annotation3](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation4](annotations) should be(None) // not found + val widgetAnnotationOpt = Annotations.findAnnotation[Widget](annotations) + widgetAnnotationOpt.isDefined should be(false) + } + + test("Annotations#secondaryConstructor 6") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[StaticSecondaryConstructorWithMethodAnnotation]) + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + annotationMap.get("widget1") should be(None) + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation2](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation3](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation4](annotations) should be(None) // not found + val widgetAnnotationOpt = Annotations.findAnnotation[Widget](annotations) + widgetAnnotationOpt.isDefined should be(false) + } + + test("Annotations#no constructor should not fail") { + // we expressly want this to no-op rather than fail outright + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations(classOf[NoConstructor]) + + annotationMap.isEmpty should be(true) + } + + test("Annotations#generic types") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations( + classOf[GenericTestCaseClass[Int]], + Seq(classOf[Object]) + ) // generic types resolve to Object + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(1) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + } + + test("Annotations#generic types 1") { + val annotationMap: Map[String, Array[Annotation]] = + Annotations.findAnnotations( + classOf[GenericTestCaseClassWithMultipleArgs[Int]], + Seq(classOf[Object], classOf[Int]) + ) // generic types resolve to Object + + annotationMap.isEmpty should be(false) + annotationMap.size should equal(2) + annotationMap("one").length should equal(1) + annotationMap("one").head.annotationType() should equal(classOf[Annotation1]) + annotationMap("two").length should equal(1) + annotationMap("two").head.annotationType() should equal(classOf[Annotation2]) + + val annotations = annotationMap.flatMap { case (_, annotations) => annotations.toList }.toArray + Annotations.findAnnotation[MarkerAnnotation](annotations) should be(None) // not found + Annotations.findAnnotation[Annotation1](annotations).isDefined should be(true) + Annotations.findAnnotation[Annotation2](annotations).isDefined should be(true) + } +} diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/BUILD.bazel b/util-reflect/src/test/scala/com/twitter/util/reflect/BUILD.bazel new file mode 100644 index 0000000000..5a39dc55a7 --- /dev/null +++ b/util-reflect/src/test/scala/com/twitter/util/reflect/BUILD.bazel @@ -0,0 +1,15 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-reflect/src/test/java/com/twitter/util/reflect", + ], +) diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/ClassesTest.scala b/util-reflect/src/test/scala/com/twitter/util/reflect/ClassesTest.scala new file mode 100644 index 0000000000..e141d75154 --- /dev/null +++ b/util-reflect/src/test/scala/com/twitter/util/reflect/ClassesTest.scala @@ -0,0 +1,100 @@ +package com.twitter.util.reflect + +import com.twitter.util.reflect.testclasses.has_underscore.ClassB +import com.twitter.util.reflect.testclasses.number_1.FooNumber +import com.twitter.util.reflect.testclasses.okNaming.Ok +import com.twitter.util.reflect.testclasses.ver2_3.{Ext, Response, This$Breaks$In$ManyWays} +import com.twitter.util.reflect.testclasses.{ClassA, Request} +import com.twitter.util.reflect.DoEverything.{DoEverything$Client, ServicePerEndpoint} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class ClassesTest extends AnyFunSuite with Matchers { + + test("Classes#simpleName") { + // these cause class.getSimpleName to blow up, ensure we don't + Classes.simpleName(classOf[Ext]) should equal("Ext") + try { + classOf[Ext].getSimpleName should equal("Ext") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + Classes.simpleName(classOf[Response]) should equal("Response") + try { + classOf[Response].getSimpleName should equal("Response") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + + // show we don't blow up + Classes.simpleName(classOf[Ok]) should equal("Ok") + try { + classOf[Ok].getSimpleName should equal("Ok") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + + // ensure we don't blow up + Classes.simpleName(classOf[ClassB]) should equal("ClassB") + try { + classOf[ClassB].getSimpleName should equal("ClassB") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + + // this causes class.getSimpleName to blow up, ensure we don't + Classes.simpleName(classOf[FooNumber]) should equal("FooNumber") + try { + classOf[FooNumber].getSimpleName should equal("FooNumber") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + + Classes.simpleName(classOf[Fungible[_]]) should equal("Fungible") + Classes.simpleName(classOf[BarService]) should equal("BarService") + Classes.simpleName(classOf[ToBarService]) should equal("ToBarService") + Classes.simpleName(classOf[java.lang.Object]) should equal("Object") + Classes.simpleName(classOf[GeneratedFooService]) should equal("GeneratedFooService") + Classes.simpleName(classOf[com.twitter.util.reflect.DoEverything[_]]) should equal( + "DoEverything") + Classes.simpleName(com.twitter.util.reflect.DoEverything.getClass) should equal("DoEverything$") + Classes.simpleName(classOf[DoEverything.ServicePerEndpoint]) should equal("ServicePerEndpoint") + Classes.simpleName(classOf[ServicePerEndpoint]) should equal("ServicePerEndpoint") + Classes.simpleName(classOf[ClassA]) should equal("ClassA") + Classes.simpleName(classOf[ClassA]) should equal(classOf[ClassA].getSimpleName) + Classes.simpleName(classOf[Request]) should equal("Request") + Classes.simpleName(classOf[Request]) should equal(classOf[Request].getSimpleName) + + Classes.simpleName(classOf[java.lang.String]) should equal("String") + Classes.simpleName(classOf[java.lang.String]) should equal( + classOf[java.lang.String].getSimpleName) + Classes.simpleName(classOf[String]) should equal("String") + Classes.simpleName(classOf[String]) should equal(classOf[String].getSimpleName) + Classes.simpleName(classOf[Int]) should equal("int") + Classes.simpleName(classOf[Int]) should equal(classOf[Int].getSimpleName) + + Classes.simpleName(classOf[DoEverything.DoEverything$Client]) should equal( + "DoEverything$Client") + Classes.simpleName(classOf[DoEverything$Client]) should equal("DoEverything$Client") + Classes.simpleName(classOf[DoEverything$Client]) should equal( + classOf[DoEverything$Client].getSimpleName) + + val manyWays = Classes.simpleName(classOf[This$Breaks$In$ManyWays]) + (manyWays == "ManyWays" /* JDK8 */ || + manyWays == "This$Breaks$In$ManyWays" /* JDK11 */ ) should be(true) + try { + classOf[This$Breaks$In$ManyWays].getSimpleName should equal("This$Breaks$In$ManyWays") + } catch { + case _: InternalError => + // do nothing -- fails in JDK8 but not JDK11 + } + } +} diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/ProxyTest.scala b/util-reflect/src/test/scala/com/twitter/util/reflect/ProxyTest.scala deleted file mode 100644 index 77a9a898a4..0000000000 --- a/util-reflect/src/test/scala/com/twitter/util/reflect/ProxyTest.scala +++ /dev/null @@ -1,241 +0,0 @@ -package com.twitter.util.reflect - -import com.twitter.util.{Future, Promise, Stopwatch} -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner - -object ProxySpec { - trait TestInterface { - def foo: String - def bar(a: Int): Option[Long] - def whoops: Int - def theVoid(): Unit - def theJavaVoid: java.lang.Void - def aFuture: Future[Int] - def aPromise: Promise[Int] - } - - class TestImpl extends TestInterface { - def foo = "foo" - def bar(a: Int) = Some(a.toLong) - def whoops = if (false) 2 else sys.error("whoops") - def theVoid() = () - def theJavaVoid = null - def aFuture = Future.value(2) - def aPromise = new Promise[Int] { - setValue(2) - } - } - - class TestClass { - val slota = "a" - val slotb = 2.0 - } -} - -@RunWith(classOf[JUnitRunner]) -class ProxyTest extends WordSpec { - - import ProxySpec._ - - "ProxyFactory" should { - "generate a factory for an interface" in { - var called = 0 - - val f1 = new ProxyFactory[TestInterface]({ m => - called += 1; m() - }) - val obj = new TestImpl - - val proxied = f1(obj) - - assert(proxied.foo == "foo") - assert(proxied.bar(2) == Some(2L)) - - assert(called == 2) - } - - "generate a factory for a class with a default constructor" in { - var called = 0 - - val f2 = new ProxyFactory[TestClass]({ m => - called += 1; m() - }) - val obj = new TestClass - - val proxied = f2(obj) - - assert(proxied.slota == "a") - assert(proxied.slotb == 2.0) - - assert(called == 2) - } - - "must not throw UndeclaredThrowableException" in { - val pf = new ProxyFactory[TestImpl](_()) - val proxied = pf(new TestImpl) - - intercept[RuntimeException] { - proxied.whoops - } - } - - "MethodCall returnsUnit must be true for unit/void methods" in { - var unitsCalled = 0 - - val pf = new ProxyFactory[TestImpl]({ call => - if (call.returnsUnit) unitsCalled += 1 - call() - }) - - val proxied = pf(new TestImpl) - proxied.foo - assert(unitsCalled == 0) - - proxied.theVoid() - assert(unitsCalled == 1) - - proxied.theJavaVoid - assert(unitsCalled == 2) - } - - "MethodCall returnsFuture must be true for methods that return a future or subclass" in { - var futuresCalled = 0 - - val pf = new ProxyFactory[TestImpl]({ call => - if (call.returnsFuture) futuresCalled += 1 - call() - }) - - val proxied = pf(new TestImpl) - - proxied.foo - assert(futuresCalled == 0) - - proxied.aFuture - assert(futuresCalled == 1) - - proxied.aPromise - assert(futuresCalled == 2) - } - - "MethodCall throws an exception when invoked without a target" in { - val pf = new ProxyFactory[TestImpl](_()) - val targetless = pf() - - intercept[NonexistentTargetException] { - targetless.foo - } - } - - "MethodCall can be invoked with alternate target" in { - val alt = new TestImpl { override def foo = "alt foo" } - val pf = new ProxyFactory[TestImpl](m => m(alt)) - val targetless = pf() - - assert(targetless.foo == "alt foo") - } - - "MethodCall can be invoked with alternate arguments" in { - val pf = new ProxyFactory[TestInterface](m => m(Array(3.asInstanceOf[AnyRef]))) - val proxied = pf(new TestImpl) - - assert(proxied.bar(2) == Some(3L)) - } - - "MethodCall can be invoked with alternate target and arguments" in { - val alt = new TestImpl { override def bar(i: Int) = Some(i.toLong * 10) } - val pf = new ProxyFactory[TestImpl](m => m(alt, Array(3.asInstanceOf[AnyRef]))) - val targetless = pf() - - assert(targetless.bar(2) == Some(30L)) - } - - // Sigh. Benchmarking has no place in unit tests. - - def time(f: => Unit) = { val elapsed = Stopwatch.start(); f; elapsed().inMilliseconds } - def benchmark(f: => Unit) = { - // warm up - for (i <- 1 to 10) f - - // measure - val results = for (i <- 1 to 7) yield time(f) - //println(results) - - // drop the top and bottom then average - results.sortWith((a, b) => a > b).slice(2, 5).sum / 3 - } - - val reflectConstructor = { () => - new ReferenceProxyFactory[TestInterface](_()) - } - val cglibConstructor = { () => - new ProxyFactory[TestInterface](_()) - } - /* - "maintains proxy creation speed" in { - val repTimes = 40000 - - val obj = new TestImpl - val factory1 = reflectConstructor() - val factory2 = cglibConstructor() - - val t1 = benchmark { for (i <- 1 to repTimes) factory1(obj) } - val t2 = benchmark { for (i <- 1 to repTimes) factory2(obj) } - - // println("proxy creation") - // println(t1) - // println(t2) - - t2 should beLessThan(200L) - } - */ - "maintains invocation speed" ignore { - val repTimes = 1500000 - - val obj = new TestImpl - val proxy1 = reflectConstructor()(obj) - val proxy2 = cglibConstructor()(obj) - - val t1 = benchmark { for (i <- 1 to repTimes) { proxy1.foo; proxy1.bar(2) } } - val t2 = benchmark { for (i <- 1 to repTimes) { proxy2.foo; proxy2.bar(2) } } - val t3 = benchmark { for (i <- 1 to repTimes) { obj.foo; obj.bar(2) } } - - // println("proxy invocation") - // println(t1) - // println(t2) - // println(t3) - - // faster than normal reflection - assert(t2 < t1) - - // less than 4x raw invocation - assert(t2 < t3 * 4) - } - } - - class ReferenceProxyFactory[I <: AnyRef: Manifest](f: (() => AnyRef) => AnyRef) { - import java.lang.reflect - - protected val interface = implicitly[Manifest[I]].runtimeClass - - private val proxyConstructor = { - reflect.Proxy - .getProxyClass(interface.getClassLoader, interface) - .getConstructor(classOf[reflect.InvocationHandler]) - } - - def apply[T <: I](instance: T) = { - proxyConstructor - .newInstance(new reflect.InvocationHandler { - def invoke(p: AnyRef, method: reflect.Method, args: Array[AnyRef]) = { - try { - f.apply(() => method.invoke(instance, args: _*)) - } catch { case e: reflect.InvocationTargetException => throw e.getCause } - } - }) - .asInstanceOf[I] - } - } -} diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/TypesTest.scala b/util-reflect/src/test/scala/com/twitter/util/reflect/TypesTest.scala new file mode 100644 index 0000000000..607188f85a --- /dev/null +++ b/util-reflect/src/test/scala/com/twitter/util/reflect/TypesTest.scala @@ -0,0 +1,238 @@ +package com.twitter.util.reflect + +import com.twitter.util.{Awaitable, FuturePool} +import com.twitter.util.reflect.testclasses._ +import com.twitter.util.reflect.testclasses.has_underscore.ClassB +import com.twitter.util.reflect.testclasses.number_1.FooNumber +import com.twitter.util.reflect.testclasses.okNaming.Ok +import com.twitter.util.reflect.testclasses.ver2_3.{Ext, Response} +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import scala.reflect.runtime.universe._ + +@RunWith(classOf[JUnitRunner]) +class TypesTest extends AnyFunSuite with Matchers { + + test("Types#asTypeTag handles classes") { + typeTagEquals(classOf[String], Types.asTypeTag(classOf[String])) + typeTagEquals(classOf[AnyRef], Types.asTypeTag(classOf[AnyRef])) + typeTagEquals(classOf[AnyVal], Types.asTypeTag(classOf[AnyVal])) + typeTagEquals(classOf[Unit], Types.asTypeTag(classOf[Unit])) + typeTagEquals(classOf[Byte], Types.asTypeTag(classOf[Byte])) + typeTagEquals(classOf[Short], Types.asTypeTag(classOf[Short])) + typeTagEquals(classOf[Char], Types.asTypeTag(classOf[Char])) + typeTagEquals(classOf[Int], Types.asTypeTag(classOf[Int])) + typeTagEquals(classOf[Long], Types.asTypeTag(classOf[Long])) + typeTagEquals(classOf[Float], Types.asTypeTag(classOf[Float])) + typeTagEquals(classOf[Double], Types.asTypeTag(classOf[Double])) + typeTagEquals(classOf[Boolean], Types.asTypeTag(classOf[Boolean])) + typeTagEquals(classOf[java.lang.Object], Types.asTypeTag(classOf[java.lang.Object])) + typeTagEquals(classOf[Any], Types.asTypeTag(classOf[Any])) + typeTagEquals(classOf[Null], Types.asTypeTag(classOf[Null])) + typeTagEquals(classOf[Nothing], Types.asTypeTag(classOf[Nothing])) + + // generics + typeTagEquals(classOf[Baz[_, _]], Types.asTypeTag(classOf[Baz[_, _]])) + typeTagEquals(classOf[Foo[_]], Types.asTypeTag(classOf[Foo[_]])) + typeTagEquals(classOf[Bez[_, _]], Types.asTypeTag(classOf[Bez[_, _]])) + typeTagEquals(classOf[TypedTrait[_, _]], Types.asTypeTag(classOf[TypedTrait[_, _]])) + } + + test("Types#runtimeClass") { + val tag = typeTag[Baz[_, _]] + Types.runtimeClass(tag) should equal(classOf[Baz[_, _]]) + Types.runtimeClass[Baz[_, _]] should equal(classOf[Baz[_, _]]) + + classOf[AnyRef].isAssignableFrom(Types.runtimeClass[AnyRef]) + classOf[AnyVal].isAssignableFrom(Types.runtimeClass[AnyVal]) + classOf[Unit].isAssignableFrom(Types.runtimeClass[Unit]) + classOf[Byte].isAssignableFrom(Types.runtimeClass[Byte]) + classOf[Short].isAssignableFrom(Types.runtimeClass[Short]) + classOf[Char].isAssignableFrom(Types.runtimeClass[Char]) + classOf[Int].isAssignableFrom(Types.runtimeClass[Int]) + classOf[Long].isAssignableFrom(Types.runtimeClass[Long]) + classOf[Float].isAssignableFrom(Types.runtimeClass[Float]) + classOf[Double].isAssignableFrom(Types.runtimeClass[Double]) + classOf[Boolean].isAssignableFrom(Types.runtimeClass[Boolean]) + classOf[java.lang.Object].isAssignableFrom(Types.runtimeClass[java.lang.Object]) + classOf[Null].isAssignableFrom(Types.runtimeClass[scala.runtime.Null$]) + classOf[Nothing].isAssignableFrom(Types.runtimeClass[scala.runtime.Nothing$]) + // `Any` is special and can result in a ClassNotFoundException + //Types.runtimeClass[Any] + } + + test("Types#equals") { + Types.equals[String](typeTag[String]) should be(true) + Types.equals[AnyRef](TypeTag.AnyRef) should be(true) + Types.equals[AnyVal](TypeTag.AnyVal) should be(true) + Types.equals[Unit](TypeTag.Unit) should be(true) + Types.equals[Byte](TypeTag.Byte) should be(true) + Types.equals[Short](TypeTag.Short) should be(true) + Types.equals[Char](TypeTag.Char) should be(true) + Types.equals[Int](TypeTag.Int) should be(true) + Types.equals[Long](TypeTag.Long) should be(true) + Types.equals[Float](TypeTag.Float) should be(true) + Types.equals[Double](TypeTag.Double) should be(true) + Types.equals[Boolean](TypeTag.Boolean) should be(true) + Types.equals[java.lang.Object](TypeTag.Object) should be(true) + Types.equals[Any](TypeTag.Any) should be(true) + Types.equals[Null](TypeTag.Null) should be(true) + Types.equals[Nothing](TypeTag.Nothing) should be(true) + // generics + Types.equals[Baz[_, _]](typeTag[Baz[_, _]]) should be(true) + Types.equals[Foo[_]](typeTag[Foo[_]]) should be(true) + Types.equals[Bez[_, _]](typeTag[Bez[_, _]]) should be(true) + Types.equals[TypedTrait[_, _]](typeTag[TypedTrait[_, _]]) should be(true) + + Types.equals[Int](TypeTag.Float) should be(false) + Types.equals[String](TypeTag.Int) should be(false) + Types.equals[java.lang.Object](TypeTag.Int) should be(false) + Types.equals[String](TypeTag.Object) should be(false) + } + + test("Types#parameterizedTypeNames") { + val clazz = classOf[Baz[_, _]] + Types.parameterizedTypeNames(clazz) should equal(Array("U", "T")) + + val constructors = clazz.getConstructors + constructors.size should equal(1) + constructors.foreach { cons => + cons.getParameters.length should equal(2) + val first = cons.getParameters.head + val second = cons.getParameters.last + Types.parameterizedTypeNames(first.getParameterizedType) should equal(Array("U")) + Types.parameterizedTypeNames(second.getParameterizedType) should equal(Array("T")) + } + + Types.parameterizedTypeNames(classOf[TypedTrait[_, _]]) should equal(Array("A", "B")) + Types.parameterizedTypeNames(classOf[Map[_, _]]) should equal(Array("K", "V")) + + val clazz2 = classOf[Bez[_, _]] + Types.parameterizedTypeNames(clazz2) should equal(Array("I", "J")) + val constructors2 = clazz2.getConstructors + constructors2.size should equal(1) + constructors2.foreach { cons => + cons.getParameters.length should equal(1) + val first = cons.getParameters.head + Types.parameterizedTypeNames(first.getParameterizedType) should equal(Array("I", "J")) + } + } + + test("Types#isCaseClass") { + Types.isCaseClass(classOf[Ext]) should be(true) + Types.isCaseClass(classOf[Response]) should be(true) + + Types.isCaseClass(classOf[Ok]) should be(true) + + Types.isCaseClass(classOf[ClassB]) should be(true) + + Types.isCaseClass(classOf[FooNumber]) should be(true) + + Types.isCaseClass(classOf[ClassA]) should be(true) + Types.isCaseClass(classOf[Request]) should be(true) + + Types.isCaseClass(classOf[Awaitable[_]]) should be(false) + Types.isCaseClass(classOf[FuturePool]) should be(false) + + Types.isCaseClass(classOf[AnyRef]) should be(false) + Types.isCaseClass(classOf[AnyVal]) should be(false) + Types.isCaseClass(classOf[Unit]) should be(false) + Types.isCaseClass(classOf[Byte]) should be(false) + Types.isCaseClass(classOf[Short]) should be(false) + Types.isCaseClass(classOf[Char]) should be(false) + Types.isCaseClass(classOf[Int]) should be(false) + Types.isCaseClass(classOf[java.lang.Integer]) should be(false) + Types.isCaseClass(classOf[Float]) should be(false) + Types.isCaseClass(classOf[Double]) should be(false) + Types.isCaseClass(classOf[Boolean]) should be(false) + Types.isCaseClass(classOf[java.lang.Object]) should be(false) + Types.isCaseClass(classOf[Any]) should be(false) + Types.isCaseClass(classOf[Null]) should be(false) + Types.isCaseClass(classOf[Nothing]) should be(false) + Types.isCaseClass(classOf[String]) should be(false) + Types.isCaseClass(classOf[java.lang.String]) should be(false) + + Types.isCaseClass(classOf[Seq[_]]) should be(false) + Types.isCaseClass(classOf[Seq[Ext]]) should be(false) + Types.isCaseClass(classOf[List[_]]) should be(false) + Types.isCaseClass(classOf[List[Ext]]) should be(false) + Types.isCaseClass(classOf[Iterable[_]]) should be(false) + Types.isCaseClass(classOf[Iterable[Ext]]) should be(false) + Types.isCaseClass(classOf[Array[_]]) should be(false) + Types.isCaseClass(classOf[Array[Ext]]) should be(false) + Types.isCaseClass(classOf[Option[_]]) should be(false) + Types.isCaseClass(classOf[Option[ClassB]]) should be(false) + Types.isCaseClass(classOf[Either[_, _]]) should be(false) + Types.isCaseClass(classOf[Either[Request, Response]]) should be(false) + Types.isCaseClass(classOf[Map[_, _]]) should be(false) + Types.isCaseClass(classOf[Map[ClassA, FooNumber]]) should be(false) + Types.isCaseClass(classOf[Tuple1[_]]) should be(false) + Types.isCaseClass(classOf[Tuple2[_, _]]) should be(false) + Types.isCaseClass(classOf[(_, _, _)]) should be(false) + Types.isCaseClass(classOf[(_, _, _, _)]) should be(false) + Types.isCaseClass(classOf[(_, _, _, _, _)]) should be(false) + Types.isCaseClass(classOf[(_, _, _, _, _, _)]) should be(false) + Types.isCaseClass(classOf[Tuple7[_, _, _, _, _, _, _]]) should be(false) + Types.isCaseClass(classOf[Tuple8[_, _, _, _, _, _, _, _]]) should be(false) + Types.isCaseClass(classOf[Tuple9[_, _, _, _, _, _, _, _, _]]) should be(false) + Types.isCaseClass(classOf[Tuple10[_, _, _, _, _, _, _, _, _, _]]) should be(false) + } + + test("Types#notCaseClass") { + Types.notCaseClass(classOf[Ext]) should be(false) + Types.notCaseClass(classOf[Response]) should be(false) + + Types.notCaseClass(classOf[Ok]) should be(false) + + Types.notCaseClass(classOf[ClassB]) should be(false) + + Types.notCaseClass(classOf[FooNumber]) should be(false) + + Types.notCaseClass(classOf[ClassA]) should be(false) + Types.notCaseClass(classOf[Request]) should be(false) + + Types.notCaseClass(classOf[Awaitable[_]]) should be(true) + Types.notCaseClass(classOf[FuturePool]) should be(true) + + Types.notCaseClass(classOf[java.lang.String]) should be(true) + Types.notCaseClass(classOf[String]) should be(true) + Types.notCaseClass(classOf[Int]) should be(true) + + Types.notCaseClass(classOf[String]) should be(true) + Types.notCaseClass(classOf[Boolean]) should be(true) + Types.notCaseClass(classOf[java.lang.Integer]) should be(true) + Types.notCaseClass(classOf[Int]) should be(true) + Types.notCaseClass(classOf[Double]) should be(true) + Types.notCaseClass(classOf[Float]) should be(true) + Types.notCaseClass(classOf[Long]) should be(true) + Types.notCaseClass(classOf[Seq[_]]) should be(true) + Types.notCaseClass(classOf[Seq[Ext]]) should be(true) + Types.notCaseClass(classOf[List[_]]) should be(true) + Types.notCaseClass(classOf[List[Ext]]) should be(true) + Types.notCaseClass(classOf[Iterable[_]]) should be(true) + Types.notCaseClass(classOf[Iterable[Ext]]) should be(true) + Types.notCaseClass(classOf[Array[_]]) should be(true) + Types.notCaseClass(classOf[Array[Ext]]) should be(true) + Types.notCaseClass(classOf[Option[_]]) should be(true) + Types.notCaseClass(classOf[Option[ClassB]]) should be(true) + Types.notCaseClass(classOf[Either[_, _]]) should be(true) + Types.notCaseClass(classOf[Either[Request, Response]]) should be(true) + Types.notCaseClass(classOf[Map[_, _]]) should be(true) + Types.notCaseClass(classOf[Map[ClassA, FooNumber]]) should be(true) + Types.notCaseClass(classOf[Tuple1[_]]) should be(true) + Types.notCaseClass(classOf[Tuple2[_, _]]) should be(true) + Types.notCaseClass(classOf[(_, _, _)]) should be(true) + Types.notCaseClass(classOf[(_, _, _, _)]) should be(true) + Types.notCaseClass(classOf[(_, _, _, _, _)]) should be(true) + Types.notCaseClass(classOf[(_, _, _, _, _, _)]) should be(true) + Types.notCaseClass(classOf[Tuple7[_, _, _, _, _, _, _]]) should be(true) + Types.notCaseClass(classOf[Tuple8[_, _, _, _, _, _, _, _]]) should be(true) + Types.notCaseClass(classOf[Tuple9[_, _, _, _, _, _, _, _, _]]) should be(true) + Types.notCaseClass(classOf[Tuple10[_, _, _, _, _, _, _, _, _, _]]) should be(true) + } + + private[this] def typeTagEquals[T](clazz: Class[T], typeTag: TypeTag[T]): Boolean = + clazz.isAssignableFrom(typeTag.mirror.runtimeClass(typeTag.tpe.typeSymbol.asClass)) +} diff --git a/util-reflect/src/test/scala/com/twitter/util/reflect/testclasses.scala b/util-reflect/src/test/scala/com/twitter/util/reflect/testclasses.scala new file mode 100644 index 0000000000..a316e73bf5 --- /dev/null +++ b/util-reflect/src/test/scala/com/twitter/util/reflect/testclasses.scala @@ -0,0 +1,130 @@ +package com.twitter.util.reflect + +import com.twitter.util.Future + +trait Fungible[+T] +trait BarService +trait ToBarService { + def toBarService: BarService +} + +abstract class GeneratedFooService + +trait DoEverything[+MM[_]] extends GeneratedFooService + +object DoEverything extends GeneratedFooService { self => + + trait ServicePerEndpoint extends ToBarService with Fungible[ServicePerEndpoint] + + class DoEverything$Client extends DoEverything[Future] +} + +object testclasses { + trait TestTraitA + trait TestTraitB + + case class Foo[T](data: T) + case class Baz[U, T](bar: U, foo: Foo[T]) + case class Bez[I, J](m: Map[I, J]) + + trait TypedTrait[A, B] + + object ver2_3 { + // this "package" naming will blow up class.getSimpleName + // packages with underscore then a number make Java unhappy + final case class Ext(a: String, b: String) + case class Response(a: String, b: String) + + class This$Breaks$In$ManyWays + } + + object okNaming { + case class Ok(a: String, b: String) + } + + object has_underscore { + case class ClassB(a: String, b: String) + } + + object number_1 { + // this "package" naming will blow up class.getSimpleName + // packages with underscore then number make Java unhappy + final case class FooNumber(a: String) + } + + final case class ClassA(a: String, b: String) + case class Request(a: String) + + class Bar + + case class NoConstructor() + + case class CaseClassOneTwo(@Annotation1 one: String, @Annotation2 two: String) + case class CaseClassOneTwoWithFields(@Annotation1 one: String, @Annotation2 two: String) { + val city: String = "San Francisco" + val state: String = "California" + } + case class CaseClassOneTwoWithAnnotatedField(@Annotation1 one: String, @Annotation2 two: String) { + @Annotation3 val three: String = "three" + } + case class CaseClassThreeFour(@Annotation3 three: String, @Annotation4 four: String) + case class CaseClassFive(@Annotation5 five: String) + case class CaseClassOneTwoThreeFour( + @Annotation1 one: String, + @Annotation2 two: String, + @Annotation3 three: String, + @Annotation4 four: String) + + case class WithThings( + @Annotation1 @Thing("thing1") thing1: String, + @Annotation2 @Thing("thing2") thing2: String) + case class WithWidgets( + @Annotation3 @Widget("widget1") widget1: String, + @Annotation4 @Widget("widget2") widget2: String) + + case class WithSecondaryConstructor( + @Annotation1 one: Int, + @Annotation2 two: Int) { + def this(@Annotation3 three: String, @Annotation4 four: String) { + this(three.toInt, four.toInt) + } + } + + object StaticSecondaryConstructor { + // NOTE: this is a factory method and not a constructor, so annotations will not be picked up + def apply(@Annotation3 three: String, @Annotation4 four: String): StaticSecondaryConstructor = + StaticSecondaryConstructor(three.toInt, four.toInt) + } + case class StaticSecondaryConstructor(@Annotation1 one: Int, @Annotation2 two: Int) + + object StaticSecondaryConstructorWithMethodAnnotation { + def apply(@Annotation3 three: String, @Annotation4 four: String): StaticSecondaryConstructor = + StaticSecondaryConstructor(three.toInt, four.toInt) + } + case class StaticSecondaryConstructorWithMethodAnnotation( + @Annotation1 one: Int, + @Annotation2 two: Int) { + // will not be found as Annotations only scans for declared field annotations, this is a method + @Widget("widget1") def widget1: String = "this is widget 1 method" + } + + case class GenericTestCaseClass[T](@Annotation1 one: T) + case class GenericTestCaseClassWithMultipleArgs[T](@Annotation1 one: T, @Annotation2 two: Int) + + trait AncestorWithAnnotations { + @Annotation1 def one: String + @Annotation2 def two: String + } + + case class CaseClassThreeFourAncestorOneTwo( + one: String, + two: String, + @Annotation3 three: String, + @Annotation4 four: String) + extends AncestorWithAnnotations + + case class CaseClassAncestorOneTwo(@Annotation5 five: String) extends AncestorWithAnnotations { + override val one: String = "one" + override val two: String = "two" + } +} diff --git a/util-registry/BUILD b/util-registry/BUILD index 0901ca475e..fcc4317f29 100644 --- a/util-registry/BUILD +++ b/util-registry/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-registry/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-registry/src/test/java", "util/util-registry/src/test/scala", diff --git a/util-registry/src/main/scala/BUILD b/util-registry/src/main/scala/BUILD index 0f12f986e2..f22c6d263a 100644 --- a/util-registry/src/main/scala/BUILD +++ b/util-registry/src/main/scala/BUILD @@ -1,12 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-registry", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "util/util-core/src/main/scala", + "util/util-registry/src/main/scala/com/twitter/util/registry", ], ) diff --git a/util-registry/src/main/scala/com/twitter/util/registry/BUILD b/util-registry/src/main/scala/com/twitter/util/registry/BUILD new file mode 100644 index 0000000000..641b4fb983 --- /dev/null +++ b/util-registry/src/main/scala/com/twitter/util/registry/BUILD @@ -0,0 +1,14 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-registry", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + ], +) diff --git a/util-registry/src/main/scala/com/twitter/util/registry/Library.scala b/util-registry/src/main/scala/com/twitter/util/registry/Library.scala index 02f32299d3..d2225d48df 100644 --- a/util-registry/src/main/scala/com/twitter/util/registry/Library.scala +++ b/util-registry/src/main/scala/com/twitter/util/registry/Library.scala @@ -16,7 +16,7 @@ object Library { * May only be called once with a given `name`. `params` and `name` must abide * by the guidelines for keys and values set in [[Registry]]. * - * @returns None if a library has already been registered with the given `name`, + * @return None if a library has already been registered with the given `name`, * or a [[Roster]] for resetting existing fields in the map otherwise. */ def register(name: String, params: Map[String, String]): Option[Roster] = { diff --git a/util-registry/src/main/scala/com/twitter/util/registry/Registry.scala b/util-registry/src/main/scala/com/twitter/util/registry/Registry.scala index 92eab2b039..006028f952 100644 --- a/util-registry/src/main/scala/com/twitter/util/registry/Registry.scala +++ b/util-registry/src/main/scala/com/twitter/util/registry/Registry.scala @@ -89,9 +89,7 @@ class SimpleRegistry extends Registry { } private[this] def sanitize(key: String): String = - key.filter { char => - char > 31 && char < 127 - } + key.filter { char => char > 31 && char < 127 } } /** diff --git a/util-registry/src/test/java/BUILD b/util-registry/src/test/java/BUILD index 57b1d46123..816cbab872 100644 --- a/util-registry/src/test/java/BUILD +++ b/util-registry/src/test/java/BUILD @@ -1,9 +1,16 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", "3rdparty/jvm/org/scala-lang:scala-library", - "util/util-registry/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-registry/src/main/scala/com/twitter/util/registry", ], ) diff --git a/util-registry/src/test/scala/BUILD b/util-registry/src/test/scala/BUILD index 0f5ffdbacf..bbf282c788 100644 --- a/util-registry/src/test/scala/BUILD +++ b/util-registry/src/test/scala/BUILD @@ -1,10 +1,18 @@ junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/mockito:mockito-core", "3rdparty/jvm/org/scalatest", - "util/util-registry/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "util/util-registry/src/main/scala/com/twitter/util/registry", ], ) diff --git a/util-registry/src/test/scala/com/twitter/util/registry/FormatterTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/FormatterTest.scala index f601cb7068..05e044d18b 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/FormatterTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/FormatterTest.scala @@ -1,12 +1,9 @@ package com.twitter.util.registry -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner import scala.util.control.NoStackTrace +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class FormatterTest extends FunSuite { +class FormatterTest extends AnyFunSuite { test("asMap generates reasonable Maps") { val registry = new SimpleRegistry registry.put(Seq("foo", "bar"), "baz") diff --git a/util-registry/src/test/scala/com/twitter/util/registry/GlobalRegistryTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/GlobalRegistryTest.scala index e14c6a9665..c8ca5345eb 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/GlobalRegistryTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/GlobalRegistryTest.scala @@ -1,9 +1,5 @@ package com.twitter.util.registry -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) class GlobalRegistryTest extends RegistryTest { def mkRegistry(): Registry = GlobalRegistry.withRegistry(new SimpleRegistry) { GlobalRegistry.get diff --git a/util-registry/src/test/scala/com/twitter/util/registry/LibraryTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/LibraryTest.scala index aad1b70b34..c6d66b4366 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/LibraryTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/LibraryTest.scala @@ -1,11 +1,8 @@ package com.twitter.util.registry -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class LibraryTest extends FunSuite { +class LibraryTest extends AnyFunSuite { test("Library.register registers libraries") { val simple = new SimpleRegistry GlobalRegistry.withRegistry(simple) { diff --git a/util-registry/src/test/scala/com/twitter/util/registry/RegistryTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/RegistryTest.scala index 1b8d15b388..bd97759b56 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/RegistryTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/RegistryTest.scala @@ -1,9 +1,9 @@ package com.twitter.util.registry import java.lang.{Character => JCharacter} -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -abstract class RegistryTest extends FunSuite { +abstract class RegistryTest extends AnyFunSuite { def mkRegistry(): Registry def name: String diff --git a/util-registry/src/test/scala/com/twitter/util/registry/RosterTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/RosterTest.scala index e70c688fcb..1bfda82cc8 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/RosterTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/RosterTest.scala @@ -2,14 +2,11 @@ package com.twitter.util.registry import java.util.logging.Logger import org.mockito.Mockito.{never, verify} -import org.mockito.Matchers.anyObject -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner -import org.scalatest.FunSuite -import org.scalatest.mockito.MockitoSugar +import org.mockito.ArgumentMatchers.any +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class RosterTest extends FunSuite with MockitoSugar { +class RosterTest extends AnyFunSuite with MockitoSugar { def withRoster(fn: (Roster, Logger) => Unit): Unit = { val simple = new SimpleRegistry simple.put(Seq("foo", "bar"), "fantastic") @@ -27,7 +24,7 @@ class RosterTest extends FunSuite with MockitoSugar { Entry(Seq("foo", "bar"), "qux") ) assert(GlobalRegistry.get.toSet == expected) - verify(log, never).warning(anyObject[String]) + verify(log, never).warning(any[String]) } } @@ -38,7 +35,7 @@ class RosterTest extends FunSuite with MockitoSugar { Entry(Seq("foo", "bar"), "fantastic") ) assert(GlobalRegistry.get.toSet == expected) - verify(log, never).warning(anyObject[String]) + verify(log, never).warning(any[String]) } } diff --git a/util-registry/src/test/scala/com/twitter/util/registry/SimpleRegistryTest.scala b/util-registry/src/test/scala/com/twitter/util/registry/SimpleRegistryTest.scala index 33a08d229f..05608e7d1c 100644 --- a/util-registry/src/test/scala/com/twitter/util/registry/SimpleRegistryTest.scala +++ b/util-registry/src/test/scala/com/twitter/util/registry/SimpleRegistryTest.scala @@ -1,9 +1,5 @@ package com.twitter.util.registry -import org.junit.runner.RunWith -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) class SimpleRegistryTest extends RegistryTest { def mkRegistry(): Registry = new SimpleRegistry def name: String = "SimpleRegistry" diff --git a/util-routing/README.md b/util-routing/README.md new file mode 100644 index 0000000000..fb0efce121 --- /dev/null +++ b/util-routing/README.md @@ -0,0 +1,5 @@ +# util-routing + +`util-routing` is a project containing base interfaces for routing an input to a destination. +This project should be considered experimental and the likelihood of API changes which result +in code breakage or removal is high. diff --git a/util-routing/src/main/scala/com/twitter/util/routing/BUILD b/util-routing/src/main/scala/com/twitter/util/routing/BUILD new file mode 100644 index 0000000000..7cddac8122 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/BUILD @@ -0,0 +1,19 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-routing", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + ], + exports = [ + "util/util-core:util-core-util", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + ], +) diff --git a/util-routing/src/main/scala/com/twitter/util/routing/ClosedRouterException.scala b/util-routing/src/main/scala/com/twitter/util/routing/ClosedRouterException.scala new file mode 100644 index 0000000000..abd38405f3 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/ClosedRouterException.scala @@ -0,0 +1,13 @@ +package com.twitter.util.routing + +import scala.util.control.NoStackTrace + +/** + * Exception thrown when an attempt to route an `Input` to a [[Router]] where `close` has been + * initiated on the [[Router]]. + * + * @param router The [[Router]] where `close` has already been initiated. + */ +case class ClosedRouterException(router: Router[_, _]) + extends RuntimeException(s"Attempting to route on a closed router with label '${router.label}'") + with NoStackTrace diff --git a/util-routing/src/main/scala/com/twitter/util/routing/Generator.scala b/util-routing/src/main/scala/com/twitter/util/routing/Generator.scala new file mode 100644 index 0000000000..8383c0c05c --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/Generator.scala @@ -0,0 +1,8 @@ +package com.twitter.util.routing + +/** + * Functional alias for creating a [[Router router]] given a label and + * all of the defined routes for the router to be generated. + */ +abstract class Generator[Input, Route, +RouterType <: Router[Input, Route]] + extends (RouterInfo[Route] => RouterType) diff --git a/util-routing/src/main/scala/com/twitter/util/routing/Result.scala b/util-routing/src/main/scala/com/twitter/util/routing/Result.scala new file mode 100644 index 0000000000..c0c9521a2d --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/Result.scala @@ -0,0 +1,35 @@ +package com.twitter.util.routing + +/** A result from attempting to find a route via a [[Router]]. */ +sealed abstract class Result + +/** + * A result that represents that the [[Router router]] could not determine a defined route + * for a given input. + */ +case object NotFound extends Result { + + /** Java-friendly Singleton accessor */ + def get(): Result = this +} + +/** + * A result that represents that the [[Router router]] could determine a defined [[Route route]] + * for a given [[Input input]]. + * + * @param input The [[Input input]] when attempting to find a route. + * @param route The [[Route route]] that was found. + * @tparam Input The input type for the [[Router]]. + * @tparam Route The route type for the [[Router]]. + */ +final case class Found[Input, Route](input: Input, route: Route) extends Result + +/** + * A result that represents that the [[Router router]] has been closed and will no longer attempt + * to find routes for any inputs. + */ +case object RouterClosed extends Result { + + /** Java-friendly Singleton accessor */ + def get(): Result = this +} diff --git a/util-routing/src/main/scala/com/twitter/util/routing/Router.scala b/util-routing/src/main/scala/com/twitter/util/routing/Router.scala new file mode 100644 index 0000000000..3244174e43 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/Router.scala @@ -0,0 +1,99 @@ +package com.twitter.util.routing + +import com.twitter.util.logging.Logging +import com.twitter.util.{Closable, ClosableOnce, Future, Time} +import java.lang.{Iterable => JIterable} +import scala.jdk.CollectionConverters._ +import scala.util.control.NonFatal + +/** + * A generic interface for routing an input to an optional matching route. + * + * @param label A label used for identifying this Router (i.e. for distinguishing between [[Router]] + * instances in error messages or for StatsReceiver scope). + * @param routes All of the `Route` routes contained within this [[Router]] + * + * This property is used to linearly determine the routes when `close()` is called. + * It may also be used as a way to determine what routes are defined within a [[Router]] + * (i.e. for generating meaningful error messages). This property does not imply + * any relationship with [[apply]] or [[find]] or that a [[Router]]'s runtime behavior + * needs to be linear. + * @tparam Input The Input type used when determining a route destination + * @tparam Route The resulting route type. A `Route` may have dynamic + * properties and may be found to satisfy multiple `Input` + * instances. As there may not be a 1-to-1 mapping of + * `Input` to `Route`, it is encouraged that a `Route` + * encapsulates its relationship to an `Input`. + * + * An example to consider would be an HTTP [[Router]]. + * An HTTP request contains a Method (i.e. GET, POST) and + * a URI. A [[Router]] for a REST API may match "GET /users/{id}" + * for multiple HTTP requests (i.e. "GET /users/123" AND "GET /users/456" would map + * to the same destination `Route`). As a result, the `Route` in this example + * should be aware of the method, path, and and the destination/logic responsible + * for handling an `Input` that matches. + * + * Another property to consider for a `Route` is "uniqueness". A [[Router]] should be + * able to determine either: + * a) a single `Route` for an `Input` + * b) no `Route` for an `Input` + * A [[Router]] should NOT be expected to return multiple routes for an `Input`. + * + * @note A [[Router]] should be considered immutable unless explicitly noted. + */ +abstract class Router[Input, +Route](val label: String, val routes: Iterable[Route]) + extends (Input => Result) + with ClosableOnce + with Logging { + + /** Java-friendly constructor */ + def this(label: String, routes: JIterable[Route]) = this(label, routes.asScala) + + /** + * Attempt to route an `Input` to a `Route` route known by this Router. + * + * @param input The `Input` to use for determining an available route + * + * @return [[Found `Found(input: Input, route: Route)`]] if a route is found to match, + * [[NotFound]] if there is no match, or + * [[RouterClosed]] if `close()` has been initiated on this [[Router]] and subsequent + * routing attempts are received. + */ + final def apply(input: Input): Result = + if (isClosed) RouterClosed + else find(input) + + /** + * Logic used to determine if this [[Router]] contains a `Route` + * for an `Input`. + * + * @param input The `Input` to determine a route for. + * + * @return [[Found `Found(input: Input, route: Route)`]] if a matching route is defined, + * [[NotFound]] if a route destination is not defined for the `Input`. + * + * @note This method is only meant to be called within the [[Router]]'s `apply`, which + * handles lifecycle concerns. It should not be accessed directly. + */ + protected def find(input: Input): Result + + /** + * @note NonFatal exceptions encountered when calling `close()` on a `Route` will be + * suppressed and Fatal exceptions will take on the same exception behavior of a `Future.join`. + */ + protected def closeOnce(deadline: Time): Future[Unit] = + Future.join(routes.map(closeIfClosable(_, deadline)).toSeq) + + private[this] def closeIfClosable(route: Route, deadline: Time): Future[Unit] = route match { + case closable: Closable => + closable.close(deadline).rescue { + case NonFatal(e) => + logger.warn( + s"Error encountered when attempting to close route '$route' in router '$label'", + e) + Future.Done + } + case _ => + Future.Done + } +} diff --git a/util-routing/src/main/scala/com/twitter/util/routing/RouterBuilder.scala b/util-routing/src/main/scala/com/twitter/util/routing/RouterBuilder.scala new file mode 100644 index 0000000000..8056c37290 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/RouterBuilder.scala @@ -0,0 +1,59 @@ +package com.twitter.util.routing + +import scala.collection.immutable.Queue + +object RouterBuilder { + + def newBuilder[Input, Route, RouterType <: Router[Input, Route]]( + generator: Generator[Input, Route, RouterType] + ): RouterBuilder[Input, Route, RouterType] = + RouterBuilder(generator) + +} + +/** + * Utility for building and creating [[Router routers]]. The resulting [[Router router]] + * should be considered immutable, unless the [[Router router's]] implementation + * explicitly states otherwise. + * + * @tparam Input The [[Router router's]] `Input` type. + * @tparam Route The [[Router router's]] destination `Route` type. It is recommended that the `Route` + * is a self-contained/self-describing type for the purpose of validation via + * the [[Validator]]. Put differently, the `Route` should know of the + * `Input` that maps to itself. + * @tparam RouterType The type of [[Router]] to build. + */ +case class RouterBuilder[Input, Route, +RouterType <: Router[Input, Route]] private ( + private val generator: Generator[Input, Route, RouterType], + private val label: String = "router", + private val routes: Queue[Route] = Queue.empty, + private val validator: Validator[Route] = Validator.None) { + + /** Set the [[Router.label label]] for the resulting [[Router router]] */ + def withLabel(label: String): RouterBuilder[Input, Route, RouterType] = + copy(label = label) + + /** + * Add the route to the routes that will be present in + * the [[Router router]] when [[newRouter]] is called. + */ + def withRoute(route: Route): RouterBuilder[Input, Route, RouterType] = + copy(routes = routes :+ route) + + /** + * Configure the route validation logic for this builder + */ + def withValidator( + validator: Validator[Route] + ): RouterBuilder[Input, Route, RouterType] = + copy(validator = validator) + + /** Generate a new [[Router router]] from the defined routes */ + def newRouter(): RouterType = { + val failures = validator(routes) + + if (failures.nonEmpty) throw ValidationException(failures) + generator(RouterInfo(label, routes)) + } + +} diff --git a/util-routing/src/main/scala/com/twitter/util/routing/RouterInfo.scala b/util-routing/src/main/scala/com/twitter/util/routing/RouterInfo.scala new file mode 100644 index 0000000000..bc89050fe3 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/RouterInfo.scala @@ -0,0 +1,14 @@ +package com.twitter.util.routing + +import java.lang.{Iterable => JIterable} +import scala.jdk.CollectionConverters._ + +/** The information needed to generate a [[Router]] via a [[Generator]]. */ +case class RouterInfo[+Route](label: String, routes: Iterable[Route]) { + + /** Java-friendly constructor. */ + def this(label: String, jRoutes: JIterable[Route]) = this(label, jRoutes.asScala) + + /** Java-friendly accessor to the `Iterable[Route]` routes. */ + def jRoutes(): JIterable[_ <: Route] = routes.asJava +} diff --git a/util-routing/src/main/scala/com/twitter/util/routing/ValidationError.scala b/util-routing/src/main/scala/com/twitter/util/routing/ValidationError.scala new file mode 100644 index 0000000000..e19e1f5a85 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/ValidationError.scala @@ -0,0 +1,9 @@ +package com.twitter.util.routing + +/** + * Container class for an error that is encountered as part of validating routes via + * a [[RouterBuilder]]. + * + * @param msg The message that explains the error. + */ +case class ValidationError(msg: String) diff --git a/util-routing/src/main/scala/com/twitter/util/routing/ValidationException.scala b/util-routing/src/main/scala/com/twitter/util/routing/ValidationException.scala new file mode 100644 index 0000000000..38bebd3d23 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/ValidationException.scala @@ -0,0 +1,17 @@ +package com.twitter.util.routing + +private object ValidationException { + def errorMsg(failures: Iterable[ValidationError]): String = failures match { + case head :: Nil => head.msg + case Nil => throw new IllegalArgumentException("RouteValidationError failures cannot be empty") + case failures => failures.map(_.msg).mkString(", ") + } +} + +/** + * Exception thrown when a [[RouterBuilder]] observes any routes that are not valid for + * the [[Router]] type it is building. + */ +case class ValidationException private[routing] (failures: Iterable[ValidationError]) + extends Exception( + s"Route Validation Failed! Errors encountered: [${ValidationException.errorMsg(failures)}]") diff --git a/util-routing/src/main/scala/com/twitter/util/routing/Validator.scala b/util-routing/src/main/scala/com/twitter/util/routing/Validator.scala new file mode 100644 index 0000000000..bf39407f71 --- /dev/null +++ b/util-routing/src/main/scala/com/twitter/util/routing/Validator.scala @@ -0,0 +1,27 @@ +package com.twitter.util.routing + +import java.lang.{Iterable => JIterable} +import java.util.function.Function +import scala.jdk.CollectionConverters._ + +object Validator { + + /** [[Validator]] that does no validation and always returns empty error results */ + val None: Validator[Any] = new Validator[Any] { + def apply(routes: Iterable[Any]): Iterable[ValidationError] = Iterable.empty + } + + /** Java-friendly creation of a [[Validator]] */ + def create[T](fn: Function[JIterable[T], JIterable[ValidationError]]): Validator[T] = + new Validator[T] { + def apply(routes: Iterable[T]): Iterable[ValidationError] = + fn(routes.asJava).asScala + } + +} + +/** + * Functional alias for determining whether all defined results are valid for a specific + * [[Router router]] type. + */ +abstract class Validator[-Route] extends (Iterable[Route] => Iterable[ValidationError]) diff --git a/util-routing/src/test/java/com/twitter/util/routing/BUILD.bazel b/util-routing/src/test/java/com/twitter/util/routing/BUILD.bazel new file mode 100644 index 0000000000..b680dc96fc --- /dev/null +++ b/util-routing/src/test/java/com/twitter/util/routing/BUILD.bazel @@ -0,0 +1,19 @@ +junit_tests( + sources = [ + "*.java", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-routing/src/main/scala/com/twitter/util/routing", + ], +) diff --git a/util-routing/src/test/java/com/twitter/util/routing/RouterCompilationTest.java b/util-routing/src/test/java/com/twitter/util/routing/RouterCompilationTest.java new file mode 100644 index 0000000000..7488298dd9 --- /dev/null +++ b/util-routing/src/test/java/com/twitter/util/routing/RouterCompilationTest.java @@ -0,0 +1,107 @@ +package com.twitter.util.routing; + +import org.junit.Assert; +import org.junit.Test; + +import java.util.Map; +import java.util.HashMap; +import java.util.stream.Collectors; +import java.util.stream.StreamSupport; + +public class RouterCompilationTest { + + private static final class StringRoute { + public final String input; + public final String output; + + StringRoute(String input, String output) { + this.input = input; + this.output = output; + } + } + + private static final class StringRouterGenerator extends Generator { + private static final StringRouterGenerator instance = new StringRouterGenerator(); + static final StringRouterGenerator getInstance() { + return instance; + } + + private StringRouterGenerator() {} + + @Override + public StringRouter apply(RouterInfo routerInfo) { + Map routeMap = new HashMap<>(); + routerInfo.jRoutes().forEach(r -> routeMap.put(r.input, r)); + return new StringRouter(routeMap); + } + } + + private static final class StringRouter extends Router { + static RouterBuilder newBuilder() { + return RouterBuilder.newBuilder(StringRouterGenerator.getInstance()); + } + + private final Map routeMap; + + public StringRouter(Map routeMap) { + super("strings", routeMap.values()); + this.routeMap = routeMap; + } + + @Override + public Result find(String input) { + StringRoute found = routeMap.get(input); + if (found == null) return NotFound.get(); + return Found.apply(input, found); + } + + } + + private static final class ValidatingStringRouterBuilder { + static RouterBuilder newBuilder() { + return StringRouter.newBuilder().withValidator(Validator.create(routes -> + StreamSupport.stream(routes.spliterator(), false) + .filter(r -> r.input.equals("invalid")) + .map(msg -> new ValidationError("INVALID!")) + .collect(Collectors.toList()))); + } + } + + @Test + public void testRouter() { + StringRoute hello = new StringRoute("hello", "abc"); + StringRoute goodbye = new StringRoute("goodbye", "123"); + + Map routes = new HashMap<>(); + routes.put(hello.input, hello); + routes.put(goodbye.input, goodbye); + + Router router = new StringRouter(routes); + Assert.assertEquals(router.apply("hello"), Found.apply("hello", hello)); + Assert.assertEquals(router.apply("goodbye"), Found.apply("goodbye", goodbye)); + Assert.assertEquals(router.apply("oh-no"), NotFound.get()); + } + + @Test + public void testRouterBuilder() { + StringRoute hello = new StringRoute("hello", "abc"); + StringRoute goodbye = new StringRoute("goodbye", "123"); + + Router router = StringRouter.newBuilder() + .withRoute(hello) + .withRoute(goodbye) + .newRouter(); + + Assert.assertEquals(router.apply("hello"), Found.apply("hello", hello)); + Assert.assertEquals(router.apply("goodbye"), Found.apply("goodbye", goodbye)); + Assert.assertEquals(router.apply("oh-no"), NotFound.get()); + } + + @Test(expected = ValidationException.class) + public void testValidatingRouterBuilder() { + ValidatingStringRouterBuilder.newBuilder() + .withRoute(new StringRoute("invalid", "other")) + .newRouter(); + } + +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/BUILD.bazel b/util-routing/src/test/scala/com/twitter/util/routing/BUILD.bazel new file mode 100644 index 0000000000..37a861211c --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/BUILD.bazel @@ -0,0 +1,19 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-routing/src/main/scala/com/twitter/util/routing", + "util/util-routing/src/test/scala/com/twitter/util/routing/dynamic", + "util/util-routing/src/test/scala/com/twitter/util/routing/simple", + "util/util-core/src/main/scala/com/twitter/conversions", + ], +) diff --git a/util-routing/src/test/scala/com/twitter/util/routing/PolymorphicRouterCompilationTest.scala b/util-routing/src/test/scala/com/twitter/util/routing/PolymorphicRouterCompilationTest.scala new file mode 100644 index 0000000000..5ffe6f8adc --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/PolymorphicRouterCompilationTest.scala @@ -0,0 +1,135 @@ +package com.twitter.util.routing + +import org.scalatest.funsuite.AnyFunSuite + +private object PolymorphicRouterCompilationTest { + + /* polymorphic output type router */ + + sealed trait Status + case object Success extends Status + case object Failure extends Status + + sealed trait StatusHandler { + def str: String + def status: Status + } + case class StatusHandlerA(str: String, status: Status) extends StatusHandler + case class StatusHandlerB(str: String, status: Status) extends StatusHandler + + object StatusRouter { + def newBuilder: RouterBuilder[String, StatusHandler, StatusRouter] = + RouterBuilder.newBuilder[String, StatusHandler, StatusRouter]( + new Generator[String, StatusHandler, StatusRouter] { + def apply(labelAndRoutes: RouterInfo[StatusHandler]): StatusRouter = + new StatusRouter(labelAndRoutes.label, labelAndRoutes.routes) + }) + } + + class StatusRouter(label: String, routes: Iterable[StatusHandler]) + extends Router[String, StatusHandler](label, routes) { + protected def find(input: String): Result = { + routes.find(_.str == input) match { + case Some(route) => Found(input, route) + case _ => NotFound + } + } + } + + val statusRouter: StatusRouter = StatusRouter.newBuilder + .withRoute(StatusHandlerA("success", Success)) + .withRoute(StatusHandlerB("fail", Failure)) + .newRouter() + + def checkStatusRouterType: Router[String, StatusHandler] = statusRouter + + val covariantStatusRouter: Router[String, StatusHandler] = + new Router[String, StatusHandlerA]("covariant-status", Iterable.empty) { + protected def find(input: String): Result = NotFound + } + + /* polymorphic input type router */ + + sealed trait Command + case class IntCommand(value: Int) extends Command + case class StringCommand(value: String) extends Command + case object NullCommand extends Command + + case class CommandHandler(canHandle: Command => Boolean, status: Status) + + object CommandRouter { + def newBuilder: RouterBuilder[Command, CommandHandler, CommandRouter] = + RouterBuilder.newBuilder(new Generator[Command, CommandHandler, CommandRouter] { + def apply(labelAndRoutes: RouterInfo[CommandHandler]): CommandRouter = + new CommandRouter(labelAndRoutes.label, labelAndRoutes.routes) + }) + } + + class CommandRouter(label: String, routes: Iterable[CommandHandler]) + extends Router[Command, CommandHandler](label, routes) { + protected def find(input: Command): Result = + routes.find(_.canHandle(input)) match { + case Some(r) => Found(input, r) + case _ => NotFound + } + } + + val CommandHandlerA: CommandHandler = CommandHandler(_.isInstanceOf[IntCommand], Success) + val CommandHandlerB: CommandHandler = CommandHandler(_.isInstanceOf[StringCommand], Failure) + + val commandRouter: CommandRouter = CommandRouter.newBuilder + .withRoute(CommandHandlerA) + .withRoute(CommandHandlerB) + .newRouter() + + def checkCommandRouterType: Router[Command, CommandHandler] = commandRouter + +} + +// these tests ensure that we properly support polymorphic Router and RouterBuilder types +class PolymorphicRouterCompilationTest extends AnyFunSuite { + import PolymorphicRouterCompilationTest._ + + test("Supports routing with polymorphic destinations") { + assert(statusRouter("success") == Found("success", StatusHandlerA("success", Success))) + assert(statusRouter("fail") == Found("fail", StatusHandlerB("fail", Failure))) + assert(statusRouter("meh") == NotFound) + } + + test("Supports routing with polymorphic inputs") { + assert(commandRouter(IntCommand(1)) == Found(IntCommand(1), CommandHandlerA)) + assert(commandRouter(StringCommand("1")) == Found(StringCommand("1"), CommandHandlerB)) + assert(commandRouter(NullCommand) == NotFound) + } + + test("Input type cannot be covariant") { + assertDoesNotCompile { + """ + | val contravariantInputRouter: Router[Command, CommandHandler] = new Router[IntCommand, CommandHandler]("int-command-router", Iterable.empty) { + | protected def find(input: IntCommand): Result = NotFound + | } + |""".stripMargin + } + } + + test("Input type cannot be contravariant") { + assertDoesNotCompile { + """ + | val contravariantInputRouter: Router[IntCommand, CommandHandler] = new Router[Command, CommandHandler]("int-command-router", Iterable.empty) { + | protected def find(input: Command): Result = NotFound + | } + |""".stripMargin + } + } + + test("Destination type cannot be contravariant") { + assertDoesNotCompile { + """ + | val contravariantStatusRouter: Router[String, StatusHandlerA] = new Router[String, StatusHandler]("contravariant-status", Iterable.empty[StatusHandler]) { + | protected def find(input: String): Result = NotFound + | } + |""".stripMargin + } + } + +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/RouterBuilderTest.scala b/util-routing/src/test/scala/com/twitter/util/routing/RouterBuilderTest.scala new file mode 100644 index 0000000000..2cfadb5a53 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/RouterBuilderTest.scala @@ -0,0 +1,90 @@ +package com.twitter.util.routing + +import com.twitter.util.routing.simple.{SimpleRoute, SimpleRouter} +import org.scalatest.funsuite.AnyFunSuite + +class RouterBuilderTest extends AnyFunSuite { + + private object TestRouter { + def newBuilder(): RouterBuilder[String, SimpleRoute, SimpleRouter] = + RouterBuilder.newBuilder(new Generator[String, SimpleRoute, SimpleRouter] { + def apply(labelAndRoutes: RouterInfo[SimpleRoute]): SimpleRouter = + new SimpleRouter(labelAndRoutes.routes.map { r => r.in -> r }.toMap) + }) + } + + private object ValidatingTestRouterBuilder { + def newBuilder(): RouterBuilder[String, SimpleRoute, SimpleRouter] = + TestRouter + .newBuilder().withValidator( + new Validator[SimpleRoute] { + def apply(routes: Iterable[SimpleRoute]): Iterable[ValidationError] = + routes.collect { + case r if r.in == "invalid" => ValidationError(s"INVALID @ $r") + } + } + ) + } + + test("can build routes") { + val router = TestRouter + .newBuilder() + .withRoute(SimpleRoute("x", true)) + .withRoute(SimpleRoute("y", false)) + .withRoute(SimpleRoute("z", true)) + .newRouter() + + assert(router("x") == Found("x", SimpleRoute("x", true))) + assert(router("y") == Found("y", SimpleRoute("y", false))) + assert(router("z") == Found("z", SimpleRoute("z", true))) + assert(router("a") == NotFound) + } + + test("throws when invalid route is passed") { + val router = ValidatingTestRouterBuilder + .newBuilder() + .withRoute(SimpleRoute("x", true)) + .withRoute(SimpleRoute("y", false)) + .withRoute(SimpleRoute("z", true)) + + router.newRouter() // everything is valid up to here + + val err = intercept[ValidationException] { + router.withRoute(SimpleRoute("invalid", false)).newRouter() + } + assert( + err.getMessage == "Route Validation Failed! Errors encountered: [INVALID @ SimpleRoute(invalid,false)]") + } + + test("throws when multiple invalid routes are passed") { + val router = ValidatingTestRouterBuilder + .newBuilder() + .withRoute(SimpleRoute("x", true)) + .withRoute(SimpleRoute("y", false)) + .withRoute(SimpleRoute("z", true)) + + router.newRouter() // everything is valid up to here + + val err = intercept[ValidationException] { + router + .withRoute(SimpleRoute("invalid", false)) + .withRoute(SimpleRoute("invalid", true)) + .newRouter() + } + assert( + err.getMessage == "Route Validation Failed! Errors encountered: [INVALID @ SimpleRoute(invalid,false), INVALID @ SimpleRoute(invalid,true)]") + } + + test("support contravariant builder") { + assertCompiles { + """ + | val typed: RouterBuilder[String, SimpleRoute, SimpleRouter] = ValidatingTestRouterBuilder.newBuilder() + | val generic: RouterBuilder[String, SimpleRoute, Router[String, SimpleRoute]] = typed + | val simpleRouter: SimpleRouter = typed.newRouter() + | val simpleRouterB: Router[String, SimpleRoute] = generic.newRouter() + | val simpleRouterC: Router[String, SimpleRoute] = typed.newRouter() + |""".stripMargin + } + } + +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/RouterTest.scala b/util-routing/src/test/scala/com/twitter/util/routing/RouterTest.scala new file mode 100644 index 0000000000..3bf19c1604 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/RouterTest.scala @@ -0,0 +1,162 @@ +package com.twitter.util.routing + +import com.twitter.conversions.DurationOps._ +import com.twitter.util.routing.dynamic.{DynamicRouter, DynamicRoute, Method, Request, Response} +import com.twitter.util.routing.simple.{SimpleRoute, SimpleRouter} +import com.twitter.util.{Await, Awaitable, Closable, Future, Time} +import java.util.concurrent.atomic.AtomicBoolean +import org.scalatest.funsuite.AnyFunSuite + +class RouterTest extends AnyFunSuite { + + private def await[T](awaitable: Awaitable[T]): T = + Await.result(awaitable, 2.seconds) + + test("can define a router of generic types") { + val helloRoute = SimpleRoute("hello", true) + val goodbyeRoute = SimpleRoute("goodbye", false) + + val router = new SimpleRouter(Map("hello" -> helloRoute, "goodbye" -> goodbyeRoute)) + + assert(router("hello") == Found("hello", helloRoute)) + assert(router("goodbye") == Found("goodbye", goodbyeRoute)) + assert(router("something else") == NotFound) + } + + test("test a router with dynamic routes") { + val abcReadHandler = + DynamicRoute(Method.Read, _.query.startsWith("abc/"), _ => Future.value(Response.Success)) + val xyzReadHandler = + DynamicRoute(Method.Read, _.query.startsWith("xyz/"), _ => Future.value(Response.Failure)) + val abcWriteHandler = + DynamicRoute(Method.Write, _.query.startsWith("abc/"), _ => Future.value(Response.Success)) + val defWriteHandler = + DynamicRoute(Method.Write, _.query == "def", _ => Future.value(Response.Failure)) + + val router: DynamicRouter = DynamicRouter + .newBuilder() + .withRoute(abcReadHandler) + .withRoute(xyzReadHandler) + .withRoute(abcWriteHandler) + .withRoute(defWriteHandler) + .newRouter() + + assert( + router(Request(Method.Read, "abc/123/456")) == Found( + Request(Method.Read, "abc/123/456"), + abcReadHandler)) + assert(router(Request(Method.Read, "def/123/456")) == NotFound) + assert( + router(Request(Method.Read, "xyz/123/456")) == Found( + Request(Method.Read, "xyz/123/456"), + xyzReadHandler)) + assert( + router(Request(Method.Write, "abc/123/456")) == Found( + Request(Method.Write, "abc/123/456"), + abcWriteHandler)) + assert( + router(Request(Method.Write, "def")) == Found(Request(Method.Write, "def"), defWriteHandler)) + assert(router(Request(Method.Write, "def/123/456")) == NotFound) + assert(router(Request(Method.Write, "xyz/123/456")) == NotFound) + + await(router.close()) + assert(router.isClosed) + router.routes.foreach(r => assert(r.isClosed)) + } + + test("multiple close() calls on a router only executes close once") { + val router = new SimpleRouter(Map.empty) + + assert(router.closedTimes.get() == 0) + await(router.close()) + assert(router.closedTimes.get() == 1) + await(router.close()) + assert(router.closedTimes.get() == 1) + await(router.close()) + assert(router.closedTimes.get() == 1) + } + + test("routing to closed router returns a RouterClosed result") { + val router = new SimpleRouter(Map.empty) + await(router.close()) + assert(router("hello") == RouterClosed) + } + + test("router closes routes that are closable") { + class AtomicClosable extends Closable { + val isClosed = new AtomicBoolean(false); + def close(deadline: Time): Future[Unit] = { + isClosed.compareAndSet(false, true) + Future.Done + } + } + + val closableA = new AtomicClosable + val closableB = new AtomicClosable + val closableC = new AtomicClosable + + class ClosableRouter + extends Router[Boolean, Closable]("closable", Seq(closableA, closableB, closableC)) { + protected def find(input: Boolean): Result = NotFound + } + + val router = new ClosableRouter + + assert(closableA.isClosed.get() == false) + assert(closableB.isClosed.get() == false) + assert(closableC.isClosed.get() == false) + + await(router.close()) + + assert(closableA.isClosed.get() == true) + assert(closableB.isClosed.get() == true) + assert(closableC.isClosed.get() == true) + + } + + test("a router is marked as closed before underlying routes complete closing") { + val closableA = Closable.make(_ => Future.never) + val closableB = Closable.make(_ => Future.Done) + + class ClosableRouter extends Router[Boolean, Closable]("closable", Seq(closableA, closableB)) { + protected def find(input: Boolean): Result = NotFound + } + + val router = new ClosableRouter + + router(true) // route doesn't throw, router is still open + + val closeFuture = router.close() // close the router, any further requests will fail + + assert(closeFuture.isDefined == false) + assert(router(true) == RouterClosed) + assert(closeFuture.isDefined == false) // the future never completes + } + + test("a router suppresses non-fatal exceptions on close") { + val closableA = Closable.make(_ => throw new IllegalStateException("BOOM")) + + class ClosableRouter extends Router[Boolean, Closable]("closable", Seq(closableA)) { + protected def find(input: Boolean): Result = NotFound + } + + val router = new ClosableRouter + await(router.close()) // close without exception + } + + test("a router bubbles up fatal exceptions on close") { + val closableA = Closable.make(_ => throw new OutOfMemoryError("BOOM")) + val closableB = Closable.make(_ => throw new OutOfMemoryError("BAM")) + + class ClosableRouter extends Router[Boolean, Closable]("closable", Seq(closableA, closableB)) { + protected def find(input: Boolean): Result = NotFound + } + + val router = new ClosableRouter + val e = intercept[OutOfMemoryError] { + await(router.close()) + } + assert(e.getMessage == "BOOM") + } + +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/BUILD.bazel b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/BUILD.bazel new file mode 100644 index 0000000000..6333be0756 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/BUILD.bazel @@ -0,0 +1,11 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + "util/util-routing/src/main/scala/com/twitter/util/routing", + ], +) diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRoute.scala b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRoute.scala new file mode 100644 index 0000000000..eb98472bb0 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRoute.scala @@ -0,0 +1,12 @@ +package com.twitter.util.routing.dynamic + +import com.twitter.util.{ClosableOnce, Future, Time} + +// example of a Route for a dynamic input +private[routing] case class DynamicRoute( + method: Method, + canHandle: Request => Boolean, + handle: Request => Future[Response]) + extends ClosableOnce { + protected def closeOnce(deadline: Time): Future[Unit] = Future.Done +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRouter.scala b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRouter.scala new file mode 100644 index 0000000000..705ff106bf --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/DynamicRouter.scala @@ -0,0 +1,37 @@ +package com.twitter.util.routing.dynamic + +import com.twitter.util.routing.{ + Found, + Generator, + RouterInfo, + NotFound, + Result, + Router, + RouterBuilder +} + +private[routing] object DynamicRouter { + def newBuilder(): RouterBuilder[Request, DynamicRoute, DynamicRouter] = + RouterBuilder.newBuilder(new Generator[Request, DynamicRoute, DynamicRouter] { + def apply(labelAndRoutes: RouterInfo[DynamicRoute]): DynamicRouter = + new DynamicRouter(labelAndRoutes.routes.toSeq) + }) +} + +// example to show how a router might be built using dynamic inputs +private[routing] class DynamicRouter(handlers: Seq[DynamicRoute]) + extends Router[Request, DynamicRoute]("dynamic-test-router", handlers) { + + // example showing how we can optimize dynamic routes to lookup by a property of the input + private[this] val handlersByMethod: Map[Method, Seq[DynamicRoute]] = handlers.groupBy(_.method) + + protected def find(input: Request): Result = + handlersByMethod.get(input.method) match { + case Some(handlers) => + handlers.find(_.canHandle(input)) match { + case Some(r) => Found(input, r) + case _ => NotFound + } + case _ => NotFound + } +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Method.scala b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Method.scala new file mode 100644 index 0000000000..1d2257d85d --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Method.scala @@ -0,0 +1,8 @@ +package com.twitter.util.routing.dynamic + +private[routing] sealed trait Method + +private[routing] object Method { + case object Read extends Method + case object Write extends Method +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Request.scala b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Request.scala new file mode 100644 index 0000000000..bdb3f2bf90 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Request.scala @@ -0,0 +1,3 @@ +package com.twitter.util.routing.dynamic + +case class Request(method: Method, query: String) diff --git a/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Response.scala b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Response.scala new file mode 100644 index 0000000000..cdc9004c38 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/dynamic/Response.scala @@ -0,0 +1,8 @@ +package com.twitter.util.routing.dynamic + +private[routing] sealed trait Response + +private[routing] object Response { + case object Success extends Response + case object Failure extends Response +} diff --git a/util-routing/src/test/scala/com/twitter/util/routing/simple/BUILD.bazel b/util-routing/src/test/scala/com/twitter/util/routing/simple/BUILD.bazel new file mode 100644 index 0000000000..31190bd87e --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/simple/BUILD.bazel @@ -0,0 +1,13 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "util/util-core:util-core-util", + "util/util-routing/src/main/scala/com/twitter/util/routing", + ], + exports = [ + ], +) diff --git a/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRoute.scala b/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRoute.scala new file mode 100644 index 0000000000..ca7c10f2df --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRoute.scala @@ -0,0 +1,3 @@ +package com.twitter.util.routing.simple + +private[routing] case class SimpleRoute(in: String, out: Boolean) diff --git a/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRouter.scala b/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRouter.scala new file mode 100644 index 0000000000..f290a2a583 --- /dev/null +++ b/util-routing/src/test/scala/com/twitter/util/routing/simple/SimpleRouter.scala @@ -0,0 +1,22 @@ +package com.twitter.util.routing.simple + +import com.twitter.util.routing.{Found, NotFound, Result, Router} +import com.twitter.util.{Future, Time} +import java.util.concurrent.atomic.AtomicInteger + +private[routing] class SimpleRouter(routeMap: Map[String, SimpleRoute]) + extends Router[String, SimpleRoute]("test-router", routeMap.values) { + + val closedTimes: AtomicInteger = new AtomicInteger(0) + + protected def find(req: String): Result = routeMap.get(req) match { + case Some(r) => Found(req, r) + case _ => NotFound + } + + override protected def closeOnce(deadline: Time): Future[Unit] = { + closedTimes.incrementAndGet() + Future.Done + } + +} diff --git a/util-security-test-certs/src/main/resources/BUILD.bazel b/util-security-test-certs/src/main/resources/BUILD.bazel new file mode 100644 index 0000000000..07015f79b6 --- /dev/null +++ b/util-security-test-certs/src/main/resources/BUILD.bazel @@ -0,0 +1,8 @@ +resources( + sources = [ + "!**/*.pyc", + "!BUILD*", + "**/*", + ], + tags = ["bazel-compatible"], +) diff --git a/util-security/src/test/resources/README b/util-security-test-certs/src/main/resources/README similarity index 100% rename from util-security/src/test/resources/README rename to util-security-test-certs/src/main/resources/README diff --git a/util-security/src/test/resources/certs/test-rsa-chain-server.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-chain-server.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-chain-server.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-chain-server.crt diff --git a/util-security/src/test/resources/certs/test-rsa-chain.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-chain.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-chain.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-chain.crt diff --git a/util-security/src/test/resources/certs/test-rsa-expired.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-expired.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-expired.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-expired.crt diff --git a/util-security-test-certs/src/main/resources/certs/test-rsa-full-cert-chain.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-full-cert-chain.crt new file mode 100644 index 0000000000..37217b6e2b --- /dev/null +++ b/util-security-test-certs/src/main/resources/certs/test-rsa-full-cert-chain.crt @@ -0,0 +1,107 @@ +-----BEGIN CERTIFICATE----- +MIIGNTCCBB2gAwIBAgICEAAwDQYJKoZIhvcNAQELBQAwgZUxCzAJBgNVBAYTAlVT +MRMwEQYDVQQIDApDYWxpZm9ybmlhMRAwDgYDVQQKDAdUd2l0dGVyMR8wHQYDVQQL +DBZDb3JlIFN5c3RlbXMgTGlicmFyaWVzMRwwGgYDVQQDDBNDU0wgSW50ZXJtZWRp +YXRlIENBMSAwHgYJKoZIhvcNAQkBFhFyeWFub0B0d2l0dGVyLmNvbTAeFw0xNzA4 +MTUxNzUyMzZaFw0yNzA4MTMxNzUyMzZaMIGvMQswCQYDVQQGEwJVUzETMBEGA1UE +CAwKQ2FsaWZvcm5pYTEWMBQGA1UEBwwNU2FuIEZyYW5jaXNjbzEQMA4GA1UECgwH +VHdpdHRlcjEfMB0GA1UECwwWQ29yZSBTeXN0ZW1zIExpYnJhcmllczEeMBwGA1UE +AwwVbG9jYWxob3N0LnR3aXR0ZXIuY29tMSAwHgYJKoZIhvcNAQkBFhFyeWFub0B0 +d2l0dGVyLmNvbTCCASIwDQYJKoZIhvcNAQEBBQADggEPADCCAQoCggEBAJ+we1c9 +H8m2p+pj3ycQozSKKmbDRQG8ajp/IgoayR1N3FqYR5CnAGRdCTAZLfQ+dMbiE45z +ascl8a/g9e02ZZOovpIYzzVbR24N+qn5jmSP0ibRuiUsNeZ/N9IZ0YjotK+qfdiE +0qE12TiWhctBpyyn9OXLwMpn+TnXi4ejTyr1xp4OwWKzgnZXVhKzZQ+pqCkA9Edj +jKCbHcgdwPDa0LFYeakGdEkoyqTcC/jfW0uzVLPdPbC2VLA3Sr+gXbpltbb378iq +btk+/OG1YCYN5rtS0JypbC5lPzuGN/6al8arTCtJ66+5w+d0OMy0MsIIINcKP6/v +u9fXXaNQ5AUA0e8CAwEAAaOCAXEwggFtMAkGA1UdEwQCMAAwEQYJYIZIAYb4QgEB +BAQDAgZAMDMGCWCGSAGG+EIBDQQmFiRPcGVuU1NMIEdlbmVyYXRlZCBTZXJ2ZXIg +Q2VydGlmaWNhdGUwHQYDVR0OBBYEFL5RSPjQgZ2VVFuBpy9dBHm0HMzeMIHTBgNV +HSMEgcswgciAFFoRgBgNkweNV0f60G8uQ5IUwoXxoYGrpIGoMIGlMQswCQYDVQQG +EwJVUzETMBEGA1UECAwKQ2FsaWZvcm5pYTEWMBQGA1UEBwwNU2FuIEZyYW5jaXNj +bzEQMA4GA1UECgwHVHdpdHRlcjEfMB0GA1UECwwWQ29yZSBTeXN0ZW1zIExpYnJh +cmllczEUMBIGA1UEAwwLQ1NMIFJvb3QgQ0ExIDAeBgkqhkiG9w0BCQEWEXJ5YW5v +QHR3aXR0ZXIuY29tggIQADAOBgNVHQ8BAf8EBAMCBaAwEwYDVR0lBAwwCgYIKwYB +BQUHAwEwDQYJKoZIhvcNAQELBQADggIBAKHlSXrMfCFE6McLvz47gfH8171+1tzB +D5FyEaG7YUqkHxN6jjLiXE/rIrq1+zTwa9YniTw+yLUnuA6ddFxNnRu4JEVYGkFV +YJNPUUsGSSIKBPPamui156BMyt4BkA1c9HrGO+L/3QE4mjJOmLQguorHlK2GGhiv +bW4r587CmPuDzNIFb6wmoU0S4xrRRVyETQE+LbdbsErZC2T8/nq8XyII7EpbzNpo +XX6CP59ajvGX11b5aAo9GmLKwQHmpBvreY8RdciW7YViNlUNdGY786tVxkUvMPwi +hsSK1ikik7Ffn7tZ3VIj0wyJFQOU08dnCrx8FTga2AuctExv1lbFL8xHwy/2ZXFB +O6df/CuIUAKyA++L8/MFyKJFb4zXQb+D8Zu5L8lmvYWhy2awfjjEtZIAfdO6BICD +nIej/4Wd/FDlTMNnE1vrKvcRPPAdWr7EPLSb3Le+O5SnQHh4VLmrPZrmsmGLp0ca +bMxyalEcrUJZr39fq6Y85c4gkn2EeYNozZZeMbhirvp3JnlTNLnhLm5v8Uo/xb8T +RAq9oGYV/QNT7gc55jP/3HXt3IHWZHZ3Uw++lg1ZKbFSm4NDmUOP/kaZREcislun +4hq72GLAGnJMEa7W0jW16hBCVDEJCRJWtDzwJXKi56DDNSoVydKLNmmqTDWswrJk +d0zC9i9vSawE +-----END CERTIFICATE----- +-----BEGIN CERTIFICATE----- +MIIGHjCCBAagAwIBAgICEAAwDQYJKoZIhvcNAQELBQAwgaUxCzAJBgNVBAYTAlVT +MRMwEQYDVQQIDApDYWxpZm9ybmlhMRYwFAYDVQQHDA1TYW4gRnJhbmNpc2NvMRAw +DgYDVQQKDAdUd2l0dGVyMR8wHQYDVQQLDBZDb3JlIFN5c3RlbXMgTGlicmFyaWVz +MRQwEgYDVQQDDAtDU0wgUm9vdCBDQTEgMB4GCSqGSIb3DQEJARYRcnlhbm9AdHdp +dHRlci5jb20wHhcNMTcwODE1MTc0ODE4WhcNMzQwMTE4MTc0ODE4WjCBlTELMAkG +A1UEBhMCVVMxEzARBgNVBAgMCkNhbGlmb3JuaWExEDAOBgNVBAoMB1R3aXR0ZXIx +HzAdBgNVBAsMFkNvcmUgU3lzdGVtcyBMaWJyYXJpZXMxHDAaBgNVBAMME0NTTCBJ +bnRlcm1lZGlhdGUgQ0ExIDAeBgkqhkiG9w0BCQEWEXJ5YW5vQHR3aXR0ZXIuY29t +MIICIjANBgkqhkiG9w0BAQEFAAOCAg8AMIICCgKCAgEA5kLopHd6fKjk3AzjGkj5 +9TnybLw85SId8MOGtEbdx88AiQA+Omidoc7efBk2c8dVu9KFrb3noQcKM6feNz9I +eul/s2hIJDlJXt3Psg/HzGH2Z46MvIc8Ebj/QtfRXDbtfVInY+xdL01o1l8o+aOD +5fQ9UE0/nPkKgCFWnQ8IAiEY9uSb7I5ujp6x68D7NUO38DhQ+6KMsU3141xb3RsV +15c6qxCIUxktNOAYv9WTkrKO4OBp64W72P2RQs6paot0001B2ARATyYYmWPWjivt +rc/eg8msBYh+xfjbqwM8XJjyvDUXwJ69TIE8KP6JwqI/elf8VPV7LAlOXwP4O8XN +fq/z/l9yNOEc6yH+25S2rA304LqD8ymJv3EtoOOERMuP/P2JtNjC5anB/96Dnjy2 +AvUY1bLVqC/A00QfW0XZw/rFd5B80rFfrdCRy6RohvaAT0EvTDZb/Q2CKKGV4KBE +aQnmWBzEjICX9OuaTWWwj8b98jCKlF00rZvJLeypoNHdjiMRmiKp+9jDbbZ2X4OS +/vz+82SgCGc/7zwAl2qR4mJa5m80JlmtFvIkipTbTfyVbxbMi6H7uXOCWyD5myAC +cNM5iO60DhbfC9yj1orO+Gy6i12A8DQmTXzVM/TLeInx3yVeoOBldK6OkLrUIc3N +VLWkwruQIYGg3cco3om2qh8CAwEAAaNmMGQwHQYDVR0OBBYEFFoRgBgNkweNV0f6 +0G8uQ5IUwoXxMB8GA1UdIwQYMBaAFH7TBrmeXnA2h6GQWlitpnUwMrGRMBIGA1Ud +EwEB/wQIMAYBAf8CAQAwDgYDVR0PAQH/BAQDAgGGMA0GCSqGSIb3DQEBCwUAA4IC +AQBHYf7aPftsjmUDRvsBxYUPncrUaY8qE1YA8JFFP7DnI2KehyITXfJd6An933IV +89SjfBb0nsDYiI36CbkFZVAJIxtYoqtpBwIMJXjgQcYI0ZHDQH3ijFrpHFYAQGm4 +2uaruekatUzrf8Mn5kgvnD72AOpqGnUXLGvxyhYK6VczDbFylNqPtJ/wDndI9xmb +Q8rbVKC6JWkllhOOROO8sG96jJuzJ9t2ob4bwfCIsXadIsHrNSDSTDH8bMqAaw8y +WprrGev80jkK2cS43w2o4hiY4L1aREGBkEja+cTJ6bb34Ebkc3Wrojtu76cHr95W +qYE0RT5/eqO04YnyPmr2/D1YTdn0z9HtJbTs7TlpxA173X7QpCG8TWki3KssBE+P +RfeHh8jVtwYT51WTZBwbf5M0bnQ1c7D3RTOAPtYkzLw65iK+ibliiqeXvXbV2FW1 +UkJCct5YuIF/NKJ0bJ75l08h86QrkS1dW5xW0IO3pN3BcVlvn+EQQRo7pKkFA6H7 +jGYbbSZP/AMZZswi9qrEp7qDt5bNjEaRQlkKLRWhOKC0GkvLEw8drbs4DcapcGwJ +vAGxcsEym1JGZQ8cQxVEq/K8Tk1blZDmJjPpplWG6YCljnvAx3Kc/BFChoprtVzg +amy+LAlsCrlMtvb9VKZShZw/IPil2d9wWQqdDOqq/Nz2gg== +-----END CERTIFICATE----- +-----BEGIN CERTIFICATE----- +MIIGMjCCBBqgAwIBAgIJANugAHtkAULIMA0GCSqGSIb3DQEBCwUAMIGlMQswCQYD +VQQGEwJVUzETMBEGA1UECAwKQ2FsaWZvcm5pYTEWMBQGA1UEBwwNU2FuIEZyYW5j +aXNjbzEQMA4GA1UECgwHVHdpdHRlcjEfMB0GA1UECwwWQ29yZSBTeXN0ZW1zIExp +YnJhcmllczEUMBIGA1UEAwwLQ1NMIFJvb3QgQ0ExIDAeBgkqhkiG9w0BCQEWEXJ5 +YW5vQHR3aXR0ZXIuY29tMB4XDTE3MDgxNTE3NDM0MloXDTM3MDgxMDE3NDM0Mlow +gaUxCzAJBgNVBAYTAlVTMRMwEQYDVQQIDApDYWxpZm9ybmlhMRYwFAYDVQQHDA1T +YW4gRnJhbmNpc2NvMRAwDgYDVQQKDAdUd2l0dGVyMR8wHQYDVQQLDBZDb3JlIFN5 +c3RlbXMgTGlicmFyaWVzMRQwEgYDVQQDDAtDU0wgUm9vdCBDQTEgMB4GCSqGSIb3 +DQEJARYRcnlhbm9AdHdpdHRlci5jb20wggIiMA0GCSqGSIb3DQEBAQUAA4ICDwAw +ggIKAoICAQC5sYrwlKIsytZbclVnrYs+KK0qY8izZ8KK4g4hQ/SRNL5YkvOecvfi +qCq19Dm5PE024Y/dAb0M4zq6HEHzsF5+M81o2FSKraiqUeJe3YwYxCwM+5msGym1 +Bys2rHnboe15Uj0KSRzW1Fouokw5GXdqTqKtmSYgHuRm0IYU8f6B2T6/tjjObUCW +Mm66xpR6YkFQtWLWgoCwAcBclvhwLyQ4GT3CBIxWzXW/AyXv09XcjgQIp2ftm2ej +bNhO/UMS2Juwzl35Sk2ilSpHGtSIwJvAwcrPFZ/LCenrC/bslbp2GMZ1JmT1+Kgt +xOUbyHRPsJGmf05kRL3+JlPhU2tBkMP5hhbr3/W+HqoSGrhGDt199HHZDy3lvESX +Hu46PsIakQj4LG58eJ98rHEa+LG4bqiROjDfOL39Fy2pkdq2Sp/nYXNS05H/uBaT +sePaWXeUAL9VcCq4sCFfeL1n7aic0Dwjv2+i1hiETvVhFzuQk1ZumNIUJBC6IOKI +nr65WWf+KQnUPQX1coOUbFSHJF/ZQAi9WqoYyLWDA3rUWVFaS6cbaF0gAsPEbYKQ +TkC1ApPxpQhCYAp8djTYzzkjsIEc6rFRGrGSHAxITtLlNZ/SLD6D4WS8n+JQMhAp +UsG0fPlEfyr/IT/MmytZgC8NgARCjpA72fMqhCXJ90RU0hq3BvPqowIDAQABo2Mw +YTAdBgNVHQ4EFgQUftMGuZ5ecDaHoZBaWK2mdTAysZEwHwYDVR0jBBgwFoAUftMG +uZ5ecDaHoZBaWK2mdTAysZEwDwYDVR0TAQH/BAUwAwEB/zAOBgNVHQ8BAf8EBAMC +AYYwDQYJKoZIhvcNAQELBQADggIBADbYfooAL6S0i+zecjZShEjKCZL9BFZAyy68 +HSgvyoOToinu9XRwLSb6OzMTQcsJ/zlETc+sb7hP0jtpIDHgNOv+Wrreh/7otKsp +/ijYGG9RfvZZQGZs0Q98/GjhqTJryXwOqL+eBX8V4SnJbOKaQgURH0eTcs2Se8r8 +4TNhwZFtwMPGbCW3HQVkx3VONVYKQAzGvbJBydj+wJVB76aO98DgsYchdgr21Mkv +elm953mnqMqWxyZzr9Ok8gsYeM23uX1OTCFMAzMxtrRTQWCRjq3NAedNVqufd4WF +eVda3NM3L1dipzyvwePh1XAQmYOxPoyKiDIKlYuumSO6CtRWT/wMRnNlEq4VHgTd +3JqeOv5gSM8Sd99/l6PdCxAJGhUPdHUOst3S4H6fy21BWd4srfivJuSRRYtWrs89 +Ij4sxaYq7PzcSTZkVV5Fn0Dkf8F5qna3AKKODbefqSssG9tw1UeavDranNiUzg3O +1LPnBurGDEWRyX+0bg60+cB6BY3mkWaBlYyFSs+8AGHgwfBn/g+rVapPA6XdJGvz +7VZXexyTe+IMWU6013JIXiwcZEUD3s3xNHJ8bdyOHf3kEz2gYMVF4q1uP8OHpE9Y +/6/0JfEjlRoG7A8mOi+QS6aYYSiCHuYLAaUL+ABbRFrLq6ppKWNuVgwsTRlqL3im +2C/dTIlS +-----END CERTIFICATE----- diff --git a/util-security/src/test/resources/certs/test-rsa-future.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-future.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-future.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-future.crt diff --git a/util-security/src/test/resources/certs/test-rsa-garbage.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-garbage.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-garbage.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-garbage.crt diff --git a/util-security/src/test/resources/certs/test-rsa-inter.crt b/util-security-test-certs/src/main/resources/certs/test-rsa-inter.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa-inter.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa-inter.crt diff --git a/util-security/src/test/resources/certs/test-rsa.crt b/util-security-test-certs/src/main/resources/certs/test-rsa.crt similarity index 100% rename from util-security/src/test/resources/certs/test-rsa.crt rename to util-security-test-certs/src/main/resources/certs/test-rsa.crt diff --git a/util-security/src/test/resources/crl/csl-intermediate-garbage.crl b/util-security-test-certs/src/main/resources/crl/csl-intermediate-garbage.crl similarity index 100% rename from util-security/src/test/resources/crl/csl-intermediate-garbage.crl rename to util-security-test-certs/src/main/resources/crl/csl-intermediate-garbage.crl diff --git a/util-security/src/test/resources/crl/csl-intermediate.crl b/util-security-test-certs/src/main/resources/crl/csl-intermediate.crl similarity index 100% rename from util-security/src/test/resources/crl/csl-intermediate.crl rename to util-security-test-certs/src/main/resources/crl/csl-intermediate.crl diff --git a/util-security/src/test/resources/keys/test-pkcs8-dsa.key b/util-security-test-certs/src/main/resources/keys/test-pkcs8-dsa.key similarity index 100% rename from util-security/src/test/resources/keys/test-pkcs8-dsa.key rename to util-security-test-certs/src/main/resources/keys/test-pkcs8-dsa.key diff --git a/util-security/src/test/resources/keys/test-pkcs8-ec.key b/util-security-test-certs/src/main/resources/keys/test-pkcs8-ec.key similarity index 100% rename from util-security/src/test/resources/keys/test-pkcs8-ec.key rename to util-security-test-certs/src/main/resources/keys/test-pkcs8-ec.key diff --git a/util-security/src/test/resources/keys/test-pkcs8-garbage.key b/util-security-test-certs/src/main/resources/keys/test-pkcs8-garbage.key similarity index 100% rename from util-security/src/test/resources/keys/test-pkcs8-garbage.key rename to util-security-test-certs/src/main/resources/keys/test-pkcs8-garbage.key diff --git a/util-security/src/test/resources/keys/test-pkcs8.key b/util-security-test-certs/src/main/resources/keys/test-pkcs8.key similarity index 100% rename from util-security/src/test/resources/keys/test-pkcs8.key rename to util-security-test-certs/src/main/resources/keys/test-pkcs8.key diff --git a/util-security/src/test/resources/pem/test-before-after.txt b/util-security-test-certs/src/main/resources/pem/test-before-after.txt similarity index 100% rename from util-security/src/test/resources/pem/test-before-after.txt rename to util-security-test-certs/src/main/resources/pem/test-before-after.txt diff --git a/util-security/src/test/resources/pem/test-multiple.txt b/util-security-test-certs/src/main/resources/pem/test-multiple.txt similarity index 100% rename from util-security/src/test/resources/pem/test-multiple.txt rename to util-security-test-certs/src/main/resources/pem/test-multiple.txt diff --git a/util-security/src/test/resources/pem/test-no-foot.txt b/util-security-test-certs/src/main/resources/pem/test-no-foot.txt similarity index 100% rename from util-security/src/test/resources/pem/test-no-foot.txt rename to util-security-test-certs/src/main/resources/pem/test-no-foot.txt diff --git a/util-security/src/test/resources/pem/test-no-head.txt b/util-security-test-certs/src/main/resources/pem/test-no-head.txt similarity index 100% rename from util-security/src/test/resources/pem/test-no-head.txt rename to util-security-test-certs/src/main/resources/pem/test-no-head.txt diff --git a/util-security/src/test/resources/pem/test.txt b/util-security-test-certs/src/main/resources/pem/test.txt similarity index 100% rename from util-security/src/test/resources/pem/test.txt rename to util-security-test-certs/src/main/resources/pem/test.txt diff --git a/util-security/BUILD b/util-security/BUILD index de525d1381..75d5e84a11 100644 --- a/util-security/BUILD +++ b/util-security/BUILD @@ -1,4 +1,5 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-security/src/main/scala", ], diff --git a/util-security/OWNERS b/util-security/OWNERS deleted file mode 100644 index cbb42ab042..0000000000 --- a/util-security/OWNERS +++ /dev/null @@ -1,7 +0,0 @@ -banderson -bdimmick -dschobel -koliver -mnakamura -roanta -ryano diff --git a/util-security/PROJECT b/util-security/PROJECT new file mode 100644 index 0000000000..291d1a5373 --- /dev/null +++ b/util-security/PROJECT @@ -0,0 +1,10 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + # Individual users are not responsible for providing a blocking review, but they + # are capable of doing so in the event that members of the groups cannot. + - jillianc + - mbezoyan + - tapans + - csl-core-security:phab + - platform-security:ldap diff --git a/util-security/src/main/scala/BUILD b/util-security/src/main/scala/BUILD index bc52977657..8596e1a998 100644 --- a/util-security/src/main/scala/BUILD +++ b/util-security/src/main/scala/BUILD @@ -1,13 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-security", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "util/util-core/src/main/scala", - "util/util-logging", + "util/util-security/src/main/scala/com/twitter/util/security", ], ) diff --git a/util-security/src/main/scala/com/twitter/util/security/BUILD b/util-security/src/main/scala/com/twitter/util/security/BUILD new file mode 100644 index 0000000000..faa70c2910 --- /dev/null +++ b/util-security/src/main/scala/com/twitter/util/security/BUILD @@ -0,0 +1,16 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-security", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/yaml:snakeyaml", + "util/util-core:util-core-util", + "util/util-logging/src/main/scala/com/twitter/logging", + ], +) diff --git a/util-security/src/main/scala/com/twitter/util/security/Credentials.scala b/util-security/src/main/scala/com/twitter/util/security/Credentials.scala new file mode 100644 index 0000000000..6bb8ceef27 --- /dev/null +++ b/util-security/src/main/scala/com/twitter/util/security/Credentials.scala @@ -0,0 +1,63 @@ +/* + * Copyright 2011 Twitter, Inc. + * + * 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 + * + * https://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 com.twitter.util.security + +import java.io.File +import org.yaml.snakeyaml.Yaml +import scala.collection.JavaConverters._ +import scala.collection.immutable.HashMap +import scala.io.Source + +/** + * Simple helper to read authentication credentials from a text file. + * + * The file's format is assumed to be yaml, containing string keys and values. + */ +object Credentials { + + def byName(name: String): Map[String, String] = { + apply(new File(sys.env.getOrElse("KEY_FOLDER", "/etc/keys"), name)) + } + + def apply(file: File): Map[String, String] = { + val fileSource = Source.fromFile(file) + try apply(fileSource.mkString) + finally fileSource.close() + } + + def apply(data: String): Map[String, String] = { + val parser = new Yaml() + val result: java.util.Map[String, Any] = parser.load(data) + if (result == null) Map.empty + else { + val builder = HashMap.newBuilder[String, String] + result.forEach { (k, v) => + builder += k -> String.valueOf(v) + } + builder.result() + } + } +} + +/** + * Java interface to Credentials. + */ +class Credentials { + def read(data: String): java.util.Map[String, String] = Credentials(data).asJava + def read(file: File): java.util.Map[String, String] = Credentials(file).asJava + def byName(name: String): java.util.Map[String, String] = Credentials.byName(name).asJava +} diff --git a/util-security/src/main/scala/com/twitter/util/security/NullSslSession.scala b/util-security/src/main/scala/com/twitter/util/security/NullSslSession.scala index 11ee7c3ff6..945a98acdf 100644 --- a/util-security/src/main/scala/com/twitter/util/security/NullSslSession.scala +++ b/util-security/src/main/scala/com/twitter/util/security/NullSslSession.scala @@ -2,7 +2,8 @@ package com.twitter.util.security import java.security.Principal import java.security.cert.Certificate -import javax.net.ssl.{SSLSession, SSLSessionContext} +import javax.net.ssl.SSLSession +import javax.net.ssl.SSLSessionContext import javax.security.cert.X509Certificate /** @@ -18,7 +19,7 @@ object NullSslSession extends SSLSession { def getLocalCertificates: Array[Certificate] = Array.empty def getLocalPrincipal: Principal = NullPrincipal def getPacketBufferSize: Int = 0 - def getPeerCertificateChain: Array[X509Certificate] = Array.empty + override def getPeerCertificateChain: Array[X509Certificate] = Array.empty def getPeerCertificates: Array[Certificate] = Array.empty def getPeerHost: String = "" def getPeerPort: Int = 0 @@ -27,8 +28,8 @@ object NullSslSession extends SSLSession { def getSessionContext: SSLSessionContext = NullSslSessionContext def getValue(name: String): Object = "" def getValueNames: Array[String] = Array.empty - def invalidate: Unit = { } + def invalidate: Unit = {} def isValid: Boolean = false - def putValue(name: String, value: Object): Unit = { } - def removeValue(name: String): Unit = { } + def putValue(name: String, value: Object): Unit = {} + def removeValue(name: String): Unit = {} } diff --git a/util-security/src/main/scala/com/twitter/util/security/NullSslSessionContext.scala b/util-security/src/main/scala/com/twitter/util/security/NullSslSessionContext.scala index 1eb5184adc..66d8f5ca2b 100644 --- a/util-security/src/main/scala/com/twitter/util/security/NullSslSessionContext.scala +++ b/util-security/src/main/scala/com/twitter/util/security/NullSslSessionContext.scala @@ -12,6 +12,6 @@ object NullSslSessionContext extends SSLSessionContext { def getSession(name: Array[Byte]): SSLSession = NullSslSession def getSessionCacheSize: Int = 0 def getSessionTimeout: Int = 0 - def setSessionCacheSize(size: Int): Unit = { } - def setSessionTimeout(seconds: Int): Unit = { } + def setSessionCacheSize(size: Int): Unit = {} + def setSessionTimeout(seconds: Int): Unit = {} } diff --git a/util-security/src/main/scala/com/twitter/util/security/PemFile.scala b/util-security/src/main/scala/com/twitter/util/security/PemBytes.scala similarity index 77% rename from util-security/src/main/scala/com/twitter/util/security/PemFile.scala rename to util-security/src/main/scala/com/twitter/util/security/PemBytes.scala index 8dae7cfce8..e5bb7e2f82 100644 --- a/util-security/src/main/scala/com/twitter/util/security/PemFile.scala +++ b/util-security/src/main/scala/com/twitter/util/security/PemBytes.scala @@ -1,8 +1,10 @@ package com.twitter.util.security import com.twitter.util.Try -import com.twitter.util.security.PemFile._ -import java.io.{BufferedReader, File, FileReader, IOException} +import com.twitter.util.security.PemBytes._ +import java.io.{BufferedReader, File, IOException, StringReader} +import java.nio.charset.StandardCharsets +import java.nio.file.Files import java.util.Base64 import scala.collection.mutable.ArrayBuffer @@ -10,8 +12,8 @@ import scala.collection.mutable.ArrayBuffer * Signals that the PEM file failed to parse properly due to a * missing boundary header or footer. */ -class InvalidPemFormatException(file: File, message: String) - extends IOException(s"PemFile (${file.getName}) failed to load: $message") +class InvalidPemFormatException(name: String, message: String) + extends IOException(s"PemBytes ($name) failed to load: $message") /** * An abstract representation of a PEM formatted file. @@ -25,10 +27,9 @@ class InvalidPemFormatException(file: File, message: String) * base64encodedbytes * -----END CERTIFICATE----- */ -class PemFile(file: File) { - - private[this] def openReader(file: File): BufferedReader = - new BufferedReader(new FileReader(file)) +class PemBytes(message: String, name: String) { + private[this] def openReader(): BufferedReader = + new BufferedReader(new StringReader(message)) private[this] def closeReader(bufReader: BufferedReader): Unit = Try(bufReader.close()) @@ -48,7 +49,7 @@ class PemFile(file: File) { private[this] def readEncodedMessage(bufReader: BufferedReader, messageType: String): String = { val contents = readEncodedMessages(bufReader, messageType) if (contents.isEmpty) - throw new InvalidPemFormatException(file, "Empty input") + throw new InvalidPemFormatException(name, "Empty input") else contents.head } @@ -66,7 +67,7 @@ class PemFile(file: File) { var line: String = bufReader.readLine() if (line == null) - throw new InvalidPemFormatException(file, "Empty input") + throw new InvalidPemFormatException(name, "Empty input") while (line != null) { if (state == PemState.Begin || state == PemState.End) { @@ -79,7 +80,7 @@ class PemFile(file: File) { } else { // PemState.Content if (line == header) { // A new matching header, but the previous one had no footer - throw new InvalidPemFormatException(file, "Missing " + footer) + throw new InvalidPemFormatException(name, "Missing " + footer) } if (line == footer) { // A good message, handle it and move to the end state @@ -95,10 +96,10 @@ class PemFile(file: File) { } if (state == PemState.Begin) { // We didn't see any matching header - throw new InvalidPemFormatException(file, "Missing " + header) + throw new InvalidPemFormatException(name, "Missing " + header) } else if (state == PemState.Content) { // We didn't see a corresponding footer for the current message - throw new InvalidPemFormatException(file, "Missing " + footer) + throw new InvalidPemFormatException(name, "Missing " + footer) } contents.toSeq @@ -115,7 +116,7 @@ class PemFile(file: File) { * only the first one will be returned. */ def readMessage(messageType: String): Try[Array[Byte]] = - Try(openReader(file)) + Try(openReader()) .flatMap(reader => Try(readDecodedMessage(reader, messageType)).ensure(closeReader(reader))) /** @@ -128,12 +129,12 @@ class PemFile(file: File) { * are ignored. */ def readMessages(messageType: String): Try[Seq[Array[Byte]]] = - Try(openReader(file)) + Try(openReader()) .flatMap(reader => Try(readDecodedMessages(reader, messageType)).ensure(closeReader(reader))) } -private object PemFile { +object PemBytes { private trait PemState private object PemState { @@ -167,4 +168,22 @@ private object PemFile { decoder.decode(encoded) } + /** + * Reads the given file and constructs a `PemBytes` with the contents. + */ + def fromFile(inputFile: File): Try[PemBytes] = { + Try { + Files.readAllBytes(inputFile.toPath) + }.flatMap { file => + fromBytes(file, inputFile.getName) + } + } + + /** + * Reads the given byte array and constructs a `PemBytes` with the contents. + */ + def fromBytes(bytes: Array[Byte], name: String): Try[PemBytes] = Try { + val buffered = new String(bytes, StandardCharsets.UTF_8) + new PemBytes(buffered, name) + } } diff --git a/util-security/src/main/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFile.scala b/util-security/src/main/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFile.scala index b9fbf33eb4..2e42f5223c 100644 --- a/util-security/src/main/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFile.scala +++ b/util-security/src/main/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFile.scala @@ -23,17 +23,17 @@ class Pkcs8EncodedKeySpecFile(file: File) { /** * Attempts to read the contents of the private key from the file. */ - def readPkcs8EncodedKeySpec(): Try[PKCS8EncodedKeySpec] = { - val pemFile = new PemFile(file) - pemFile - .readMessage(MessageType) - .map(new PKCS8EncodedKeySpec(_)) - .onFailure(logException) - } + def readPkcs8EncodedKeySpec(): Try[PKCS8EncodedKeySpec] = + PemBytes.fromFile(file).flatMap { pemBytes => + pemBytes + .readMessage(MessageType) + .map(new PKCS8EncodedKeySpec(_)) + .onFailure(logException) + } } -private object Pkcs8EncodedKeySpecFile { - private val MessageType: String = "PRIVATE KEY" +object Pkcs8EncodedKeySpecFile { + val MessageType: String = "PRIVATE KEY" private val log = Logger.get("com.twitter.util.security") } diff --git a/util-security/src/main/scala/com/twitter/util/security/Pkcs8KeyManagerFactory.scala b/util-security/src/main/scala/com/twitter/util/security/Pkcs8KeyManagerFactory.scala index 3b2629a230..44d4d9f6e2 100644 --- a/util-security/src/main/scala/com/twitter/util/security/Pkcs8KeyManagerFactory.scala +++ b/util-security/src/main/scala/com/twitter/util/security/Pkcs8KeyManagerFactory.scala @@ -11,56 +11,38 @@ import java.util.UUID import javax.net.ssl.{KeyManager, KeyManagerFactory} /** - * A factory which can create a [[javax.net.ssl.KeyManager KeyManager]] which contains - * an X.509 Certificate and a PKCS#8 private key. + * A factory which can create a `javax.net.ssl.KeyManager` which contains + * an X.509 Certificate (or Certificate chain) and a PKCS#8 private key. */ -class Pkcs8KeyManagerFactory(certFile: File, keyFile: File) { +class Pkcs8KeyManagerFactory(certsFile: File, keyFile: File) { private[this] def logException(ex: Throwable): Unit = log.warning( - s"Pkcs8KeyManagerFactory (${certFile.getName()}, ${keyFile.getName()}) " + + s"Pkcs8KeyManagerFactory (${certsFile.getName()}, ${keyFile.getName()}) " + s"failed to create key manager: ${ex.getMessage()}." ) - private[this] def keySpecToPrivateKey(keySpec: PKCS8EncodedKeySpec): PrivateKey = { - val kf: KeyFactory = KeyFactory.getInstance("RSA") - kf.generatePrivate(keySpec) - } - - private[this] def createKeyStore(cert: X509Certificate, privateKey: PrivateKey): KeyStore = { - val alias: String = UUID.randomUUID().toString() - val ks: KeyStore = KeyStore.getInstance("JKS") - ks.load(null) - ks.setKeyEntry(alias, privateKey, "".toCharArray(), Array(cert)) - ks - } - - private[this] def keyStoreToKeyManagers(keyStore: KeyStore): Array[KeyManager] = { - val kmf = KeyManagerFactory.getInstance("SunX509") - kmf.init(keyStore, "".toCharArray()) - kmf.getKeyManagers() - } - /** - * Attempts to read the contents of both the X.509 Certificate file and the PKCS#8 - * Private Key file and combine the contents into a [[javax.net.ssl.KeyManager KeyManager]]. + * Attempts to read the contents of both the X.509 Certificates file and the PKCS#8 + * Private Key file and combine the contents into a `javax.net.ssl.KeyManager` * The singular value is returned in an Array for ease of use with - * [[javax.net.ssl.SSLContext SSLContext's]] init method. + * `javax.net.ssl.SSLContext` init method. */ def getKeyManagers(): Try[Array[KeyManager]] = { - val tryCert: Try[X509Certificate] = new X509CertificateFile(certFile).readX509Certificate() + val tryCerts: Try[Seq[X509Certificate]] = + new X509CertificateFile(certsFile).readX509Certificates() val tryKeySpec: Try[PKCS8EncodedKeySpec] = new Pkcs8EncodedKeySpecFile(keyFile).readPkcs8EncodedKeySpec() val tryPrivateKey: Try[PrivateKey] = tryKeySpec.map(keySpecToPrivateKey) - val tryCertKey: Try[(X509Certificate, PrivateKey)] = join(tryCert, tryPrivateKey) - val tryKeyStore: Try[KeyStore] = tryCertKey.map((createKeyStore _).tupled) - tryKeyStore.map(keyStoreToKeyManagers).onFailure(logException) + val tryCertsKey: Try[(Seq[X509Certificate], PrivateKey)] = join(tryCerts, tryPrivateKey) + val tryKeyStore: Try[KeyStore] = tryCertsKey.map((createKeyStore _).tupled) + tryKeyStore.map(keyStoreToKeyManagerFactory(_).getKeyManagers()).onFailure(logException) } } -private object Pkcs8KeyManagerFactory { +object Pkcs8KeyManagerFactory { private val log = Logger.get("com.twitter.util.security") private def join[A, B](tryA: Try[A], tryB: Try[B]): Try[(A, B)] = { @@ -70,4 +52,36 @@ private object Pkcs8KeyManagerFactory { case (_, Throw(_)) => tryB.asInstanceOf[Try[(A, B)]] } } + + /** + * Convert a `PKCS8EncodedKeySpec` into a `java.security.PrivateKey` + */ + def keySpecToPrivateKey(keySpec: PKCS8EncodedKeySpec): PrivateKey = { + val kf: KeyFactory = KeyFactory.getInstance("RSA") + kf.generatePrivate(keySpec) + } + + /** + * Use the loaded `java.security.KeyStore` to construct a + * `javax.net.ssl.KeyManagerFactory`. + */ + def keyStoreToKeyManagerFactory(keyStore: KeyStore): KeyManagerFactory = { + val kmf = KeyManagerFactory.getInstance("SunX509") + kmf.init(keyStore, "".toCharArray()) + kmf + } + + /** + * Load passed in X.509 certificates into a `java.security.KeyStore` + */ + def createKeyStore( + certs: Seq[X509Certificate], + privateKey: PrivateKey + ): KeyStore = { + val alias: String = UUID.randomUUID().toString() + val ks: KeyStore = KeyStore.getInstance("JKS") + ks.load(null) + ks.setKeyEntry(alias, privateKey, "".toCharArray(), certs.toArray) + ks + } } diff --git a/util-security/src/main/scala/com/twitter/util/security/PrivateKeyFile.scala b/util-security/src/main/scala/com/twitter/util/security/PrivateKeyFile.scala index 744b550cfe..b588983955 100644 --- a/util-security/src/main/scala/com/twitter/util/security/PrivateKeyFile.scala +++ b/util-security/src/main/scala/com/twitter/util/security/PrivateKeyFile.scala @@ -21,6 +21,28 @@ class PrivateKeyFile(file: File) { private[this] def logException(ex: Throwable): Unit = log.warning(s"PrivateKeyFile (${file.getName()}) failed to load: ${ex.getMessage()}.") + /** + * Attemps to read the contents of an unencrypted PrivateKey file. + * + * @note Throws `InvalidPemFormatException` if the file is not a PEM format file. + * + * @note Throws `InvalidKeySpecException` if the file cannot be loaded as an + * RSA, DSA, or EC PrivateKey. + * + * @note This method will attempt to load the key first as an RSA key, then + * as a DSA key, and then as an EC (Elliptic Curve) key. + */ + def readPrivateKey(): Try[PrivateKey] = + new Pkcs8EncodedKeySpecFile(file) + .readPkcs8EncodedKeySpec() + .flatMap(keySpecToPrivateKey) + .onFailure(logException) + +} + +object PrivateKeyFile { + private val log = Logger.get("com.twitter.util.security") + /** * Attempts to convert a `PKCS8EncodedKeySpec` to a `PrivateKey`. * @@ -28,7 +50,7 @@ class PrivateKeyFile(file: File) { * keeps identical behavior to netty * https://github.com/netty/netty/blob/netty-4.1.11.Final/handler/src/main/java/io/netty/handler/ssl/SslContext.java#L1006 */ - private[this] def keySpecToPrivateKey(keySpec: PKCS8EncodedKeySpec): Try[PrivateKey] = + def keySpecToPrivateKey(keySpec: PKCS8EncodedKeySpec): Try[PrivateKey] = Try { KeyFactory.getInstance("RSA").generatePrivate(keySpec) }.handle { @@ -43,27 +65,4 @@ class PrivateKeyFile(file: File) { case ex: InvalidKeySpecException => Throw(new InvalidKeySpecException("None of RSA, DSA, EC worked", ex)) } - - /** - * Attemps to read the contents of an unencrypted PrivateKey file. - * - * @throws InvalidPemFormatException if the file is not a PEM format - * file. - * - * @throws InvalidKeySpecException if the file cannot be loaded as an - * RSA, DSA, or EC PrivateKey. - * - * @note This method will attempt to load the key first as an RSA key, then - * as a DSA key, and then as an EC (Elliptic Curve) key. - */ - def readPrivateKey(): Try[PrivateKey] = - new Pkcs8EncodedKeySpecFile(file) - .readPkcs8EncodedKeySpec() - .flatMap(keySpecToPrivateKey) - .onFailure(logException) - -} - -private object PrivateKeyFile { - private val log = Logger.get("com.twitter.util.security") } diff --git a/util-security/src/main/scala/com/twitter/util/security/X500PrincipalInfo.scala b/util-security/src/main/scala/com/twitter/util/security/X500PrincipalInfo.scala index 588bc6a79c..4835e5f632 100644 --- a/util-security/src/main/scala/com/twitter/util/security/X500PrincipalInfo.scala +++ b/util-security/src/main/scala/com/twitter/util/security/X500PrincipalInfo.scala @@ -3,10 +3,10 @@ package com.twitter.util.security import javax.naming.ldap.{LdapName, Rdn} import javax.security.auth.x500.X500Principal import com.twitter.util.security.X500PrincipalInfo._ -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ /** - * Parses an [[javax.security.auth.x500.X500Principal X500Principal]] into + * Parses an `javax.security.auth.x500.X500Principal` into * its individual pieces for easily extracting widely used items like an X.509 * certificate's Common Name. */ diff --git a/util-security/src/main/scala/com/twitter/util/security/X509CertificateDeserializer.scala b/util-security/src/main/scala/com/twitter/util/security/X509CertificateDeserializer.scala new file mode 100644 index 0000000000..48ab081d3e --- /dev/null +++ b/util-security/src/main/scala/com/twitter/util/security/X509CertificateDeserializer.scala @@ -0,0 +1,79 @@ +package com.twitter.util.security + +import com.twitter.util.Try + +import java.io.ByteArrayInputStream +import java.security.cert.CertificateFactory +import java.security.cert.X509Certificate + +/** + * A helper object to deserialize PEM-encoded X.509 Certificates. + * + * @example + * -----BEGIN CERTIFICATE----- + * base64encodedbytes + * -----END CERTIFICATE----- + */ +object X509CertificateDeserializer { + private[this] val MessageType: String = "CERTIFICATE" + private[this] val deserializeX509: Array[Byte] => X509Certificate = { (certBytes: Array[Byte]) => + val certFactory = CertificateFactory.getInstance("X.509") + val certificate = certFactory + .generateCertificate(new ByteArrayInputStream(certBytes)) + .asInstanceOf[X509Certificate] + certificate.checkValidity() + certificate + } + + /** + * Deserializes an [[InputStream]] that contains a PEM-encoded X.509 + * Certificate. + * + * Closes the InputStream once it has finished reading. + */ + def deserializeCertificate(rawPem: String, name: String): Try[X509Certificate] = { + val pemBytes = new PemBytes(rawPem, name) + val message: Try[Array[Byte]] = pemBytes + .readMessage(MessageType) + + message.map(deserializeX509) + } + + /** + * Deserializes an [[InputStream]] that contains a PEM-encoded X.509 + * Certificate. + * + * Closes the InputStream once it has finished reading. + */ + def deserializeCertificates(rawPem: String, name: String): Try[Seq[X509Certificate]] = { + val pemBytes = new PemBytes(rawPem, name) + val messages: Try[Seq[Array[Byte]]] = pemBytes + .readMessages(MessageType) + + messages.map(_.map(deserializeX509)) + } + + /** + * Deserializes an [[InputStream]] that contains PEM-encoded X.509 + * Certificates. Wraps the `deserializeX509` call in a Try + * (as `certificate.checkValidity()` can return CertificateExpiredException, CertificateNotYetValidException) + * and separates out any expired or not yet valid certificates detected. + * + * Closes the InputStream once it has finished reading. + */ + def deserializeAndFilterOutInvalidCertificates( + rawPem: String, + name: String + ): (Seq[Try[X509Certificate]], Seq[Try[X509Certificate]]) = { + val pemBytes = new PemBytes(rawPem, name) + val messages: Try[Seq[Array[Byte]]] = pemBytes + .readMessages(MessageType) + messages + .map(certs => { + certs + .map(cert => { + Try(deserializeX509(cert)) + }).partition(_.isReturn) + }).get() + } +} diff --git a/util-security/src/main/scala/com/twitter/util/security/X509CertificateFile.scala b/util-security/src/main/scala/com/twitter/util/security/X509CertificateFile.scala index 683be5439b..26a8b6ff38 100644 --- a/util-security/src/main/scala/com/twitter/util/security/X509CertificateFile.scala +++ b/util-security/src/main/scala/com/twitter/util/security/X509CertificateFile.scala @@ -3,8 +3,10 @@ package com.twitter.util.security import com.twitter.logging.Logger import com.twitter.util.Try import com.twitter.util.security.X509CertificateFile._ -import java.io.{ByteArrayInputStream, File} -import java.security.cert.{CertificateFactory, X509Certificate} +import java.io.File +import java.nio.charset.StandardCharsets +import java.nio.file.Files +import java.security.cert.X509Certificate /** * A representation of an X.509 Certificate PEM-encoded and stored @@ -23,38 +25,23 @@ class X509CertificateFile(file: File) { /** * Attempts to read the contents of the X.509 Certificate from the file. */ - def readX509Certificate(): Try[X509Certificate] = { - val pemFile = new PemFile(file) - pemFile - .readMessage(MessageType) - .map(generateX509Certificate) - .onFailure(logException) - } + def readX509Certificate(): Try[X509Certificate] = Try { + new String(Files.readAllBytes(file.toPath), StandardCharsets.UTF_8) + }.flatMap { buffered => + X509CertificateDeserializer.deserializeCertificate(buffered, file.getName) + }.onFailure(logException) /** * Attempts to read the contents of multiple X.509 Certificates from the file. */ - def readX509Certificates(): Try[Seq[X509Certificate]] = { - val pemFile = new PemFile(file) - pemFile - .readMessages(MessageType) - .map(certBytes => certBytes.map(generateX509Certificate)) - .onFailure(logException) - } + def readX509Certificates(): Try[Seq[X509Certificate]] = Try { + new String(Files.readAllBytes(file.toPath), StandardCharsets.UTF_8) + }.flatMap { buffered => + X509CertificateDeserializer.deserializeCertificates(buffered, file.getName) + }.onFailure(logException) } private object X509CertificateFile { - private val MessageType: String = "CERTIFICATE" - private val log = Logger.get("com.twitter.util.security") - - private def generateX509Certificate(decodedMessage: Array[Byte]): X509Certificate = { - val certFactory = CertificateFactory.getInstance("X.509") - val certificate = certFactory - .generateCertificate(new ByteArrayInputStream(decodedMessage)) - .asInstanceOf[X509Certificate] - certificate.checkValidity() - certificate - } } diff --git a/util-security/src/main/scala/com/twitter/util/security/X509CrlFile.scala b/util-security/src/main/scala/com/twitter/util/security/X509CrlFile.scala index 9816ffc8e6..dc965c6efe 100644 --- a/util-security/src/main/scala/com/twitter/util/security/X509CrlFile.scala +++ b/util-security/src/main/scala/com/twitter/util/security/X509CrlFile.scala @@ -9,7 +9,7 @@ import java.security.cert.{CertificateFactory, X509CRL} /** * A representation of an X.509 Certificate Revocation List (CRL) * PEM-encoded and stored in a file. - * + * * @example * -----BEGIN X509 CRL----- * base64encodedbytes @@ -20,13 +20,13 @@ class X509CrlFile(file: File) { private[this] def logException(ex: Throwable): Unit = log.warning(s"X509Crl (${file.getName}) failed to load: ${ex.getMessage}.") - def readX509Crl(): Try[X509CRL] = { - val pemFile = new PemFile(file) - pemFile - .readMessage(MessageType) - .map(generateX509Crl) - .onFailure(logException) - } + def readX509Crl(): Try[X509CRL] = + PemBytes.fromFile(file).flatMap { pemBytes => + pemBytes + .readMessage(MessageType) + .map(generateX509Crl) + .onFailure(logException) + } } diff --git a/util-security/src/main/scala/com/twitter/util/security/X509TrustManagerFactory.scala b/util-security/src/main/scala/com/twitter/util/security/X509TrustManagerFactory.scala index 31ae19a267..0e13c2f824 100644 --- a/util-security/src/main/scala/com/twitter/util/security/X509TrustManagerFactory.scala +++ b/util-security/src/main/scala/com/twitter/util/security/X509TrustManagerFactory.scala @@ -10,7 +10,7 @@ import java.util.UUID import javax.net.ssl.{TrustManager, TrustManagerFactory} /** - * A factory which can create a [[javax.net.ssl.TrustManager TrustManager]] which contains + * A factory which can create a `javax.net.ssl.TrustManager` which contains * a collection of X.509 certificates. */ class X509TrustManagerFactory(certsFile: File) { @@ -21,44 +21,56 @@ class X509TrustManagerFactory(certsFile: File) { s"failed to create trust manager: ${ex.getMessage()}." ) + /** + * Attempts to read the contents of a file containing a collection of X.509 certificates, + * and combines them into a `javax.net.ssl.TrustManager`. + * The singular value is returned in an Array for ease of use with + * `javax.net.ssl.SSLContext`'s init method. + */ + def getTrustManagers(): Try[Array[TrustManager]] = { + val tryCerts: Try[Seq[X509Certificate]] = + new X509CertificateFile(certsFile).readX509Certificates() + tryCerts.flatMap(buildTrustManager).onFailure(logException) + } + +} + +object X509TrustManagerFactory { + private val log = Logger.get("com.twitter.util.security") + private[this] def setCertificateEntry(ks: KeyStore)(cert: X509Certificate): Unit = { val alias: String = UUID.randomUUID().toString() ks.setCertificateEntry(alias, cert) } - private[this] def certsToKeyStore(certs: Seq[X509Certificate]): KeyStore = { + /** + * Load passed in X.509 certificates into a `java.security.KeyStore` + */ + def certsToKeyStore(certs: Seq[X509Certificate]): KeyStore = { val ks: KeyStore = KeyStore.getInstance("JKS") ks.load(null) certs.foreach(setCertificateEntry(ks)) ks } - private[this] def keyStoreToTmf(ks: KeyStore): TrustManagerFactory = { + /** + * Combine the passed in X.509 certificates in order to construct a + * `javax.net.ssl.TrustManagerFactory`. + */ + def certsToTrustManagerFactory(certs: Seq[X509Certificate]): TrustManagerFactory = { + val ks: KeyStore = certsToKeyStore(certs) val tmf = TrustManagerFactory.getInstance(TrustManagerFactory.getDefaultAlgorithm()) tmf.init(ks) tmf } - private[this] def certsToTrustManagers(certs: Seq[X509Certificate]): Array[TrustManager] = { - val ks: KeyStore = certsToKeyStore(certs) - val tmf: TrustManagerFactory = keyStoreToTmf(ks) - tmf.getTrustManagers() - } - /** - * Attempts to read the contents of a file containing a collection of X.509 certificates, - * and combines them into a [[javax.net.ssl.TrustManager TrustManager]]. + * Attempts to combine the passed in X.509 certificates in order to construct a + * `javax.net.ssl.TrustManager`. * The singular value is returned in an Array for ease of use with - * [[javax.net.ssl.SSLContext SSLContext's]] init method. + * `javax.net.ssl.SSLContext`'s init method. */ - def getTrustManagers(): Try[Array[TrustManager]] = { - val tryCerts: Try[Seq[X509Certificate]] = - new X509CertificateFile(certsFile).readX509Certificates() - tryCerts.map(certsToTrustManagers).onFailure(logException) + def buildTrustManager(x509Certificates: Seq[X509Certificate]): Try[Array[TrustManager]] = Try { + certsToTrustManagerFactory(x509Certificates).getTrustManagers } - -} - -private object X509TrustManagerFactory { - private val log = Logger.get("com.twitter.util.security") } diff --git a/util-security/src/test/java/BUILD b/util-security/src/test/java/BUILD deleted file mode 100644 index 7f674e6c57..0000000000 --- a/util-security/src/test/java/BUILD +++ /dev/null @@ -1,9 +0,0 @@ -junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "util/util-core/src/main/scala", - "util/util-security/src/main/scala", - ], -) diff --git a/util-security/src/test/java/BUILD.bazel b/util-security/src/test/java/BUILD.bazel new file mode 100644 index 0000000000..da97a848c8 --- /dev/null +++ b/util-security/src/test/java/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-security/src/test/java/com/twitter/util/security", + ], +) diff --git a/util-security/src/test/java/com/twitter/util/security/BUILD.bazel b/util-security/src/test/java/com/twitter/util/security/BUILD.bazel new file mode 100644 index 0000000000..1a9ce8496b --- /dev/null +++ b/util-security/src/test/java/com/twitter/util/security/BUILD.bazel @@ -0,0 +1,16 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-security/src/main/scala/com/twitter/util/security", + ], +) diff --git a/util-security/src/test/resources/BUILD b/util-security/src/test/resources/BUILD deleted file mode 100644 index e0afa86adb..0000000000 --- a/util-security/src/test/resources/BUILD +++ /dev/null @@ -1,6 +0,0 @@ -resources( - sources = rglobs( - "*", - exclude = [rglobs("BUILD*", "*.pyc")], - ), -) diff --git a/util-security/src/test/scala/BUILD b/util-security/src/test/scala/BUILD deleted file mode 100644 index e5abb760cf..0000000000 --- a/util-security/src/test/scala/BUILD +++ /dev/null @@ -1,12 +0,0 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-logging", - "util/util-security/src/main/scala", - "util/util-security/src/test/resources", - ], -) diff --git a/util-security/src/test/scala/BUILD.bazel b/util-security/src/test/scala/BUILD.bazel new file mode 100644 index 0000000000..019705b791 --- /dev/null +++ b/util-security/src/test/scala/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-security/src/test/scala/com/twitter/util/security", + ], +) diff --git a/util-security/src/test/scala/com/twitter/util/security/BUILD.bazel b/util-security/src/test/scala/com/twitter/util/security/BUILD.bazel new file mode 100644 index 0000000000..cadc34321e --- /dev/null +++ b/util-security/src/test/scala/com/twitter/util/security/BUILD.bazel @@ -0,0 +1,21 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-security-test-certs/src/main/resources", + "util/util-security/src/main/scala/com/twitter/util/security", + ], +) diff --git a/util-core/src/test/scala/com/twitter/util/CredentialsTest.scala b/util-security/src/test/scala/com/twitter/util/security/CredentialsTest.scala similarity index 57% rename from util-core/src/test/scala/com/twitter/util/CredentialsTest.scala rename to util-security/src/test/scala/com/twitter/util/security/CredentialsTest.scala index 35f95e3c70..46ad0da301 100644 --- a/util-core/src/test/scala/com/twitter/util/CredentialsTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/CredentialsTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -14,57 +14,55 @@ * limitations under the License. */ -package com.twitter.util +package com.twitter.util.security -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.Checkers +import org.scalatestplus.scalacheck.Checkers +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class CredentialsTest extends FunSuite with Checkers { +class CredentialsTest extends AnyFunSuite with Checkers { test("parse a simple auth file") { - val content = "username: root\npassword: hellokitty\n" - assert(Credentials(content) == Map("username" -> "root", "password" -> "hellokitty")) + val content = "username: root\npassword: not_a_password\n" + val result = Credentials(content) + assert(result == Map("username" -> "root", "password" -> "not_a_password")) } test("parse a more complex auth file") { val content = """# random comment -username:root -password : last_0f-the/international:playboys +username: root +password: not_a-real/pr0d:password # more stuff - moar :ok +moar: ok """ assert( Credentials(content) == Map( "username" -> "root", - "password" -> "last_0f-the/international:playboys", + "password" -> "not_a-real/pr0d:password", "moar" -> "ok" ) ) } test("work for java peeps too") { - val content = "username: root\npassword: hellokitty\n" + val content = "username: root\npassword: not_a_password\n" val jmap = new Credentials().read(content) assert(jmap.size() == 2) assert(jmap.get("username") == "root") - assert(jmap.get("password") == "hellokitty") + assert(jmap.get("password") == "not_a_password") } test("handle \r\n line breaks") { - val content = "username: root\r\npassword: hellokitty\r\n" - assert(Credentials(content) == Map("username" -> "root", "password" -> "hellokitty")) + val content = "username: root\r\npassword: not_a_password\r\n" + assert(Credentials(content) == Map("username" -> "root", "password" -> "not_a_password")) } test("handle special chars") { - val pass = (0 to 127) + val pass = (32 to 126) .map(_.toChar) - .filter(c => c != '\r' && c != '\n') + .filter(c => c != '\r' && c != '\n' && c != '\'') .mkString - val content = s"username: root\npassword: $pass\n" + val content = s"username: root\npassword: '$pass'\n" assert(Credentials(content) == Map("username" -> "root", "password" -> pass)) } } diff --git a/util-security/src/test/scala/com/twitter/util/security/NullPrincipalTest.scala b/util-security/src/test/scala/com/twitter/util/security/NullPrincipalTest.scala index 8ea6c82e72..ba208b8f06 100644 --- a/util-security/src/test/scala/com/twitter/util/security/NullPrincipalTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/NullPrincipalTest.scala @@ -1,9 +1,9 @@ package com.twitter.util.security import java.security.Principal -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class NullPrincipalTest extends FunSuite { +class NullPrincipalTest extends AnyFunSuite { test("NullPrincipal is a Principal") { val principal: Principal = NullPrincipal diff --git a/util-security/src/test/scala/com/twitter/util/security/NullSslSessionContextTest.scala b/util-security/src/test/scala/com/twitter/util/security/NullSslSessionContextTest.scala index c86d64d892..0a22dab61d 100644 --- a/util-security/src/test/scala/com/twitter/util/security/NullSslSessionContextTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/NullSslSessionContextTest.scala @@ -1,9 +1,9 @@ package com.twitter.util.security import javax.net.ssl.SSLSessionContext -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class NullSslSessionContextTest extends FunSuite { +class NullSslSessionContextTest extends AnyFunSuite { test("NullSslSessionContext is an SSLSessionContext") { val sslSessionContext: SSLSessionContext = NullSslSessionContext diff --git a/util-security/src/test/scala/com/twitter/util/security/NullSslSessionTest.scala b/util-security/src/test/scala/com/twitter/util/security/NullSslSessionTest.scala index 233891974b..bc2d5027fb 100644 --- a/util-security/src/test/scala/com/twitter/util/security/NullSslSessionTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/NullSslSessionTest.scala @@ -1,9 +1,9 @@ package com.twitter.util.security import javax.net.ssl.SSLSession -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class NullSslSessionTest extends FunSuite { +class NullSslSessionTest extends AnyFunSuite { test("NullSslSession is an SSLSession") { val sslSession: SSLSession = NullSslSession @@ -44,4 +44,3 @@ class NullSslSessionTest extends FunSuite { } } - diff --git a/util-security/src/test/scala/com/twitter/util/security/PemBytesTest.scala b/util-security/src/test/scala/com/twitter/util/security/PemBytesTest.scala new file mode 100644 index 0000000000..4f1c94849b --- /dev/null +++ b/util-security/src/test/scala/com/twitter/util/security/PemBytesTest.scala @@ -0,0 +1,144 @@ +package com.twitter.util.security + +import com.twitter.io.TempFile +import com.twitter.util.Try +import java.io.File +import java.nio.file.Files +import org.scalatest.funsuite.AnyFunSuite + +class PemBytesTest extends AnyFunSuite { + + private[this] def readPemBytesMessageFromBytes( + messageType: String + )( + tempBytes: Array[Byte] + ): Try[String] = + for { + pemBytes <- PemBytes.fromBytes(tempBytes, "pem") + item <- pemBytes.readMessage(messageType) + } yield new String(item) + + private[this] def readPemBytesMessagesFromBytes( + messageType: String + )( + tempBytes: Array[Byte] + ): Try[Seq[String]] = + for { + pemBytes <- PemBytes.fromBytes(tempBytes, "pem") + items <- pemBytes.readMessages(messageType) + } yield items.map(new String(_)) + + private[this] def readPemBytesMessage(messageType: String)(tempFile: File): Try[String] = + for { + pemBytes <- PemBytes.fromFile(tempFile) + item <- pemBytes.readMessage(messageType) + } yield new String(item) + + private[this] def readPemBytesMessages(messageType: String)(tempFile: File): Try[Seq[String]] = + for { + pemBytes <- PemBytes.fromFile(tempFile) + items <- pemBytes.readMessages(messageType) + } yield items.map(new String(_)) + + private[this] val readHello = readPemBytesMessage("hello") _ + private[this] val readHellos = readPemBytesMessages("hello") _ + private[this] val readHellofromBytes = readPemBytesMessageFromBytes("hello") _ + private[this] val readHellosfromBytes = readPemBytesMessagesFromBytes("hello") _ + + test("no header file") { + val tempFile = TempFile.fromResourcePath("/pem/test-no-head.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHello(tempFile) + PemBytesTestUtils.assertException[InvalidPemFormatException, String](tryHello) + PemBytesTestUtils.assertExceptionMessageContains("Missing -----BEGIN HELLO-----")(tryHello) + } + + test("no footer file") { + val tempFile = TempFile.fromResourcePath("/pem/test-no-foot.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHello(tempFile) + PemBytesTestUtils.assertException[InvalidPemFormatException, String](tryHello) + PemBytesTestUtils.assertExceptionMessageContains("Missing -----END HELLO-----")(tryHello) + } + + test("wrong message type") { + val tempFile = TempFile.fromResourcePath("/pem/test.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readPemBytesMessage("GOODBYE")(tempFile) + PemBytesTestUtils.assertException[InvalidPemFormatException, String](tryHello) + PemBytesTestUtils.assertExceptionMessageContains("Missing -----BEGIN GOODBYE-----")(tryHello) + } + + test("good file") { + val tempFile = TempFile.fromResourcePath("/pem/test.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHello(tempFile) + assert(tryHello.isReturn) + val hello = tryHello.get() + assert(hello == "hello") + } + + test("good bytes data") { + val tempBytes = Files.readAllBytes(TempFile.fromResourcePath("/pem/test.txt").toPath) + // deleteOnExit is handled by TempFile + + val tryHello = readHellofromBytes(tempBytes) + assert(tryHello.isReturn) + val hello = tryHello.get() + assert(hello == "hello") + } + + test("good file with text before and after") { + val tempFile = TempFile.fromResourcePath("/pem/test-before-after.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHello(tempFile) + assert(tryHello.isReturn) + val hello = tryHello.get() + assert(hello == "hello") + } + + test("read single message from good file with multiple messages") { + val tempFile = TempFile.fromResourcePath("/pem/test-multiple.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHello(tempFile) + assert(tryHello.isReturn) + val hello = tryHello.get() + assert(hello == "hello") + } + + test("read multiple message from good file with multiple messages") { + val tempFile = TempFile.fromResourcePath("/pem/test-multiple.txt") + // deleteOnExit is handled by TempFile + + val tryHello = readHellos(tempFile) + assert(tryHello.isReturn) + val messages = tryHello.get() + assert(messages == Seq("hello", "goodbye")) + } + + test("read single message from good encoded data with multiple messages as bytes") { + val tempBytes = Files.readAllBytes(TempFile.fromResourcePath("/pem/test-multiple.txt").toPath) + // deleteOnExit is handled by TempFile + + val tryHello = readHellofromBytes(tempBytes) + assert(tryHello.isReturn) + val hello = tryHello.get() + assert(hello == "hello") + } + + test("read multiple message from good encoded data with multiple messages") { + val tempBytes = Files.readAllBytes(TempFile.fromResourcePath("/pem/test-multiple.txt").toPath) + // deleteOnExit is handled by TempFile + + val tryHello = readHellosfromBytes(tempBytes) + assert(tryHello.isReturn) + val messages = tryHello.get() + assert(messages == Seq("hello", "goodbye")) + } +} diff --git a/util-security/src/test/scala/com/twitter/util/security/PemFileTestUtils.scala b/util-security/src/test/scala/com/twitter/util/security/PemBytesTestUtils.scala similarity index 53% rename from util-security/src/test/scala/com/twitter/util/security/PemFileTestUtils.scala rename to util-security/src/test/scala/com/twitter/util/security/PemBytesTestUtils.scala index ac6b55a41c..e183991517 100644 --- a/util-security/src/test/scala/com/twitter/util/security/PemFileTestUtils.scala +++ b/util-security/src/test/scala/com/twitter/util/security/PemBytesTestUtils.scala @@ -1,27 +1,11 @@ package com.twitter.util.security -import com.twitter.logging.{BareFormatter, Logger, StringHandler} import com.twitter.util.Try import java.io.{File, IOException} import org.scalatest.Assertions._ import scala.reflect.ClassTag -object PemFileTestUtils { - - def newLogger(): Logger = { - val logger = Logger.get("com.twitter.util.security") - logger.clearHandlers() - logger.setLevel(Logger.WARNING) - logger.setUseParentHandlers(false) - logger - } - - def newHandler(): StringHandler = { - val logger = newLogger() - val handler = new StringHandler(BareFormatter, None) - logger.addHandler(handler) - handler - } +object PemBytesTestUtils { def assertException[A <: AnyRef, B](tryValue: Try[B])(implicit classTag: ClassTag[A]): Unit = { assert(tryValue.isThrow) @@ -42,20 +26,7 @@ object PemFileTestUtils { def assertIOException[A](tryValue: Try[A]): Unit = assertException[IOException, A](tryValue) - def assertLogMessage(name: String)( - logMessage: String, - filename: String, - expectedExMessage: String - ): Unit = { - assert(logMessage.startsWith(s"$name ($filename) failed to load: ")) - assert(logMessage.contains(expectedExMessage)) - } - - def testFileDoesntExist[A]( - name: String, - read: File => Try[A] - ): Unit = { - val handler = newHandler() + def testFileDoesntExist[A](name: String, read: File => Try[A]): Unit = { val tempFile = File.createTempFile("test", "ext") tempFile.deleteOnExit() @@ -64,30 +35,20 @@ object PemFileTestUtils { assert(!tempFile.exists()) val tryReadPem: Try[A] = read(tempFile) - assertLogMessage(name)(handler.get, tempFile.getName(), "No such file or directory") assertIOException(tryReadPem) } - def testFilePathIsntFile[A]( - name: String, - read: File => Try[A] - ): Unit = { - val handler = newHandler() + def testFilePathIsntFile[A](name: String, read: File => Try[A]): Unit = { val tempFile = File.createTempFile("test", "ext") tempFile.deleteOnExit() val parentFile = tempFile.getParentFile() val tryReadPem: Try[A] = read(parentFile) - assertLogMessage(name)(handler.get, parentFile.getName(), "Is a directory") assertIOException(tryReadPem) } - def testFilePathIsntReadable[A]( - name: String, - read: File => Try[A] - ): Unit = { - val handler = newHandler() + def testFilePathIsntReadable[A](name: String, read: File => Try[A]): Unit = { val tempFile = File.createTempFile("test", "ext") tempFile.deleteOnExit() @@ -97,21 +58,20 @@ object PemFileTestUtils { val tryReadPem: Try[A] = read(tempFile) - assertLogMessage(name)(handler.get, tempFile.getName(), "Permission denied") assertIOException(tryReadPem) } def testEmptyFile[A <: AnyRef, B]( name: String, read: File => Try[B] - )(implicit classTag: ClassTag[A]): Unit = { - val handler = newHandler() + )( + implicit classTag: ClassTag[A] + ): Unit = { val tempFile = File.createTempFile("test", "ext") tempFile.deleteOnExit() val tryReadPem: Try[B] = read(tempFile) - assertLogMessage(name)(handler.get, tempFile.getName(), "Empty input") assertException[A, B](tryReadPem) } diff --git a/util-security/src/test/scala/com/twitter/util/security/PemFileTest.scala b/util-security/src/test/scala/com/twitter/util/security/PemFileTest.scala deleted file mode 100644 index f8a5cdae61..0000000000 --- a/util-security/src/test/scala/com/twitter/util/security/PemFileTest.scala +++ /dev/null @@ -1,95 +0,0 @@ -package com.twitter.util.security - -import com.twitter.io.TempFile -import com.twitter.util.Try -import java.io.File -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class PemFileTest extends FunSuite { - - private[this] val assertLogMessage = - PemFileTestUtils.assertLogMessage("PemFile") _ - - private[this] def readPemFileMessage(messageType: String)(tempFile: File): Try[String] = { - val pemFile = new PemFile(tempFile) - pemFile.readMessage(messageType).map(new String(_)) - } - - private[this] def readPemFileMessages(messageType: String)(tempFile: File): Try[Seq[String]] = { - val pemFile = new PemFile(tempFile) - pemFile.readMessages(messageType).map(items => items.map(new String(_))) - } - - private[this] val readHello = readPemFileMessage("hello") _ - private[this] val readHellos = readPemFileMessages("hello") _ - - test("no header file") { - val tempFile = TempFile.fromResourcePath("/pem/test-no-head.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHello(tempFile) - PemFileTestUtils.assertException[InvalidPemFormatException, String](tryHello) - PemFileTestUtils.assertExceptionMessageContains("Missing -----BEGIN HELLO-----")(tryHello) - } - - test("no footer file") { - val tempFile = TempFile.fromResourcePath("/pem/test-no-foot.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHello(tempFile) - PemFileTestUtils.assertException[InvalidPemFormatException, String](tryHello) - PemFileTestUtils.assertExceptionMessageContains("Missing -----END HELLO-----")(tryHello) - } - - test("wrong message type") { - val tempFile = TempFile.fromResourcePath("/pem/test.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readPemFileMessage("GOODBYE")(tempFile) - PemFileTestUtils.assertException[InvalidPemFormatException, String](tryHello) - PemFileTestUtils.assertExceptionMessageContains("Missing -----BEGIN GOODBYE-----")(tryHello) - } - - test("good file") { - val tempFile = TempFile.fromResourcePath("/pem/test.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHello(tempFile) - assert(tryHello.isReturn) - val hello = tryHello.get() - assert(hello == "hello") - } - - test("good file with text before and after") { - val tempFile = TempFile.fromResourcePath("/pem/test-before-after.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHello(tempFile) - assert(tryHello.isReturn) - val hello = tryHello.get() - assert(hello == "hello") - } - - test("read single message from good file with multiple messages") { - val tempFile = TempFile.fromResourcePath("/pem/test-multiple.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHello(tempFile) - assert(tryHello.isReturn) - val hello = tryHello.get() - assert(hello == "hello") - } - - test("read multiple message from good file with multiple messages") { - val tempFile = TempFile.fromResourcePath("/pem/test-multiple.txt") - // deleteOnExit is handled by TempFile - - val tryHello = readHellos(tempFile) - assert(tryHello.isReturn) - val messages = tryHello.get() - assert(messages == Seq("hello", "goodbye")) - } -} diff --git a/util-security/src/test/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFileTest.scala b/util-security/src/test/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFileTest.scala index 29a1a4acb5..0f2feb14de 100644 --- a/util-security/src/test/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFileTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/Pkcs8EncodedKeySpecFileTest.scala @@ -4,15 +4,9 @@ import com.twitter.io.TempFile import com.twitter.util.Try import java.io.File import java.security.spec.PKCS8EncodedKeySpec -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class Pkcs8EncodedKeySpecFileTest extends FunSuite { - - private[this] val assertLogMessage = - PemFileTestUtils.assertLogMessage("PKCS8EncodedKeySpec") _ +class Pkcs8EncodedKeySpecFileTest extends AnyFunSuite { private[this] val readKeySpecFromFile: File => Try[PKCS8EncodedKeySpec] = (tempFile) => { @@ -21,19 +15,19 @@ class Pkcs8EncodedKeySpecFileTest extends FunSuite { } test("File path doesn't exist") { - PemFileTestUtils.testFileDoesntExist("PKCS8EncodedKeySpec", readKeySpecFromFile) + PemBytesTestUtils.testFileDoesntExist("PKCS8EncodedKeySpec", readKeySpecFromFile) } test("File path isn't a file") { - PemFileTestUtils.testFilePathIsntFile("PKCS8EncodedKeySpec", readKeySpecFromFile) + PemBytesTestUtils.testFilePathIsntFile("PKCS8EncodedKeySpec", readKeySpecFromFile) } test("File path isn't readable") { - PemFileTestUtils.testFilePathIsntReadable("PKCS8EncodedKeySpec", readKeySpecFromFile) + PemBytesTestUtils.testFilePathIsntReadable("PKCS8EncodedKeySpec", readKeySpecFromFile) } test("File isn't a key spec") { - PemFileTestUtils.testEmptyFile[InvalidPemFormatException, PKCS8EncodedKeySpec]( + PemBytesTestUtils.testEmptyFile[InvalidPemFormatException, PKCS8EncodedKeySpec]( "PKCS8EncodedKeySpec", readKeySpecFromFile ) diff --git a/util-security/src/test/scala/com/twitter/util/security/Pkcs8KeyManagerFactoryTest.scala b/util-security/src/test/scala/com/twitter/util/security/Pkcs8KeyManagerFactoryTest.scala index 232795eee4..5b5f3726fa 100644 --- a/util-security/src/test/scala/com/twitter/util/security/Pkcs8KeyManagerFactoryTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/Pkcs8KeyManagerFactoryTest.scala @@ -3,14 +3,11 @@ package com.twitter.util.security import com.twitter.io.TempFile import java.io.File import javax.net.ssl.{KeyManager, X509KeyManager} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class PKCS8KeyManagerFactoryTest extends FunSuite { +class PKCS8KeyManagerFactoryTest extends AnyFunSuite { - private[this] def assertContainsCertAndKey(keyManager: KeyManager): Unit = { + private[this] def assertContainsCertsAndKey(keyManager: KeyManager, numCerts: Int = 1): Unit = { assert(keyManager.isInstanceOf[X509KeyManager]) val x509km = keyManager.asInstanceOf[X509KeyManager] @@ -22,7 +19,7 @@ class PKCS8KeyManagerFactoryTest extends FunSuite { val alias = clientAliases.head val certChain = x509km.getCertificateChain(alias) - assert(certChain.length == 1) + assert(certChain.length == numCerts) val cert = certChain.head assert(cert.getSubjectX500Principal.getName().contains("localhost.twitter.com")) @@ -83,7 +80,23 @@ class PKCS8KeyManagerFactoryTest extends FunSuite { assert(tryKeyManagers.isReturn) val keyManagers = tryKeyManagers.get() assert(keyManagers.length == 1) - assertContainsCertAndKey(keyManagers.head) + assertContainsCertsAndKey(keyManagers.head) + } + + test("Chain of certificates is good") { + val tempCertsFile = TempFile.fromResourcePath("/certs/test-rsa-full-cert-chain.crt") + // deleteOnExit is handled by TempFile + + val tempKeyFile = TempFile.fromResourcePath("/keys/test-pkcs8.key") + // deleteOnExit is handled by TempFile + + val factory = new Pkcs8KeyManagerFactory(tempCertsFile, tempKeyFile) + val tryKeyManagers = factory.getKeyManagers() + + assert(tryKeyManagers.isReturn) + val keyManagers = tryKeyManagers.get() + assert(keyManagers.length == 1) + assertContainsCertsAndKey(keyManagers.head, numCerts = 3) } } diff --git a/util-security/src/test/scala/com/twitter/util/security/PrivateKeyFileTest.scala b/util-security/src/test/scala/com/twitter/util/security/PrivateKeyFileTest.scala index f7e5e25f0c..48f7f1ac4b 100644 --- a/util-security/src/test/scala/com/twitter/util/security/PrivateKeyFileTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/PrivateKeyFileTest.scala @@ -2,9 +2,9 @@ package com.twitter.util.security import com.twitter.io.TempFile import java.security.spec.InvalidKeySpecException -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class PrivateKeyFileTest extends FunSuite { +class PrivateKeyFileTest extends AnyFunSuite { test("File is garbage") { // Lines were manually deleted from a real pkcs 8 pem file diff --git a/util-security/src/test/scala/com/twitter/util/security/X500PrincipalInfoTest.scala b/util-security/src/test/scala/com/twitter/util/security/X500PrincipalInfoTest.scala index bebdb2b406..530aa6bf23 100644 --- a/util-security/src/test/scala/com/twitter/util/security/X500PrincipalInfoTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/X500PrincipalInfoTest.scala @@ -1,12 +1,9 @@ package com.twitter.util.security import javax.security.auth.x500.X500Principal -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class X500PrincipalInfoTest extends FunSuite { +class X500PrincipalInfoTest extends AnyFunSuite { test("Empty principal") { val principal = new X500Principal("") diff --git a/util-security/src/test/scala/com/twitter/util/security/X509CertificateDeserializerTest.scala b/util-security/src/test/scala/com/twitter/util/security/X509CertificateDeserializerTest.scala new file mode 100644 index 0000000000..1d62afc58b --- /dev/null +++ b/util-security/src/test/scala/com/twitter/util/security/X509CertificateDeserializerTest.scala @@ -0,0 +1,76 @@ +package com.twitter.util.security + +import java.security.cert.{CertificateException, X509Certificate} +import org.scalatest.funsuite.AnyFunSuite +import scala.io.Source + +class X509CertificateDeserializerTest extends AnyFunSuite { + private[this] def getResourceFileAsString(path: String): String = { + val resourceStream = getClass.getResourceAsStream(path) + val resourceSource = Source.fromInputStream(resourceStream) + try { + val resource = resourceSource.mkString + resource + } finally { + resourceSource.close() + } + } + + private[this] def assertIsCslCert(cert: X509Certificate): Unit = { + val subjectData: String = cert.getSubjectX500Principal.getName() + assert(subjectData.contains("Core Systems Libraries")) + } + + test("Input is garbage") { + // Lines were manually deleted from a real certificate file + val input = getResourceFileAsString("/certs/test-rsa-garbage.crt") + + val tryCert = X509CertificateDeserializer.deserializeCertificate(input, "test") + + PemBytesTestUtils.assertException[CertificateException, X509Certificate](tryCert) + } + + test("Input is an X509 Certificate") { + val input = getResourceFileAsString("/certs/test-rsa.crt") + + val tryCert = X509CertificateDeserializer.deserializeCertificate(input, "test") + + assert(tryCert.isReturn) + val cert = tryCert.get() + + assert(cert.getSigAlgName == "SHA256withRSA") + } + + test("Input with multiple X509 Certificates") { + val input = getResourceFileAsString("/certs/test-rsa-chain.crt") + + val tryCerts = X509CertificateDeserializer.deserializeCertificates(input, "test") + + assert(tryCerts.isReturn) + val certs = tryCerts.get() + + assert(certs.length == 2) + + val intermediate = certs.head + assertIsCslCert(intermediate) + + val root = certs(1) + assertIsCslCert(root) + } + + test("X509 Certificate File is not valid: expired") { + val input = getResourceFileAsString("/certs/test-rsa-expired.crt") + + intercept[java.security.cert.CertificateExpiredException] { + X509CertificateDeserializer.deserializeCertificate(input, "test").get() + } + } + + test("X509 Certificate File is not valid: not yet ready") { + val input = getResourceFileAsString("/certs/test-rsa-future.crt") + + intercept[java.security.cert.CertificateNotYetValidException] { + X509CertificateDeserializer.deserializeCertificate(input, "test").get() + } + } +} diff --git a/util-security/src/test/scala/com/twitter/util/security/X509CertificateFileTest.scala b/util-security/src/test/scala/com/twitter/util/security/X509CertificateFileTest.scala index 377fb7bd57..d88be71db9 100644 --- a/util-security/src/test/scala/com/twitter/util/security/X509CertificateFileTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/X509CertificateFileTest.scala @@ -4,23 +4,17 @@ import com.twitter.io.TempFile import com.twitter.util.Try import java.io.File import java.security.cert.{CertificateException, X509Certificate} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class X509CertificateFileTest extends FunSuite { +class X509CertificateFileTest extends AnyFunSuite { private[this] def assertIsCslCert(cert: X509Certificate): Unit = { val subjectData: String = cert.getSubjectX500Principal.getName() assert(subjectData.contains("Core Systems Libraries")) } - private[this] val assertLogMessage = - PemFileTestUtils.assertLogMessage("X509Certificate") _ - private[this] def assertCertException(tryCert: Try[X509Certificate]): Unit = - PemFileTestUtils.assertException[CertificateException, X509Certificate](tryCert) + PemBytesTestUtils.assertException[CertificateException, X509Certificate](tryCert) private[this] val readCertFromFile: File => Try[X509Certificate] = (tempFile) => { @@ -29,26 +23,25 @@ class X509CertificateFileTest extends FunSuite { } test("File path doesn't exist") { - PemFileTestUtils.testFileDoesntExist("X509Certificate", readCertFromFile) + PemBytesTestUtils.testFileDoesntExist("X509Certificate", readCertFromFile) } test("File path isn't a file") { - PemFileTestUtils.testFilePathIsntFile("X509Certificate", readCertFromFile) + PemBytesTestUtils.testFilePathIsntFile("X509Certificate", readCertFromFile) } test("File path isn't readable") { - PemFileTestUtils.testFilePathIsntReadable("X509Certificate", readCertFromFile) + PemBytesTestUtils.testFilePathIsntReadable("X509Certificate", readCertFromFile) } test("File isn't a certificate") { - PemFileTestUtils.testEmptyFile[InvalidPemFormatException, X509Certificate]( + PemBytesTestUtils.testEmptyFile[InvalidPemFormatException, X509Certificate]( "X509Certificate", readCertFromFile ) } test("File is garbage") { - val handler = PemFileTestUtils.newHandler() // Lines were manually deleted from a real certificate file val tempFile = TempFile.fromResourcePath("/certs/test-rsa-garbage.crt") // deleteOnExit is handled by TempFile @@ -56,7 +49,6 @@ class X509CertificateFileTest extends FunSuite { val certFile = new X509CertificateFile(tempFile) val tryCert = certFile.readX509Certificate() - assertLogMessage(handler.get, tempFile.getName, "Incomplete BER/DER data.") assertCertException(tryCert) } @@ -91,23 +83,4 @@ class X509CertificateFileTest extends FunSuite { val root = certs(1) assertIsCslCert(root) } - - test("X509 Certificate File is not valid: expired") { - val tempFile = TempFile.fromResourcePath("/certs/test-rsa-expired.crt") - // deleteOnExit is handled by TempFile - - intercept[java.security.cert.CertificateExpiredException] { - readCertFromFile(tempFile).get() - } - } - - test("X509 Certificate File is not valid: not yet ready") { - val tempFile = TempFile.fromResourcePath("/certs/test-rsa-future.crt") - // deleteOnExit is handled by TempFile - - intercept[java.security.cert.CertificateNotYetValidException] { - readCertFromFile(tempFile).get() - } - } - } diff --git a/util-security/src/test/scala/com/twitter/util/security/X509CrlFileTest.scala b/util-security/src/test/scala/com/twitter/util/security/X509CrlFileTest.scala index 91ed6814e9..227b2819c8 100644 --- a/util-security/src/test/scala/com/twitter/util/security/X509CrlFileTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/X509CrlFileTest.scala @@ -4,15 +4,12 @@ import com.twitter.io.TempFile import com.twitter.util.Try import java.io.File import java.security.cert.{CRLException, X509CRL} -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class X509CrlFileTest extends FunSuite { - - private[this] val assertLogMessage = - PemFileTestUtils.assertLogMessage("X509Crl") _ +class X509CrlFileTest extends AnyFunSuite { private[this] def assertCrlException(tryCrl: Try[X509CRL]): Unit = - PemFileTestUtils.assertException[CRLException, X509CRL](tryCrl) + PemBytesTestUtils.assertException[CRLException, X509CRL](tryCrl) private[this] val readCrlFromFile: File => Try[X509CRL] = (tempFile) => { @@ -21,37 +18,35 @@ class X509CrlFileTest extends FunSuite { } test("File path doesn't exist") { - PemFileTestUtils.testFileDoesntExist("X509Crl", readCrlFromFile) + PemBytesTestUtils.testFileDoesntExist("X509Crl", readCrlFromFile) } test("File path isn't a file") { - PemFileTestUtils.testFilePathIsntFile("X509Crl", readCrlFromFile) + PemBytesTestUtils.testFilePathIsntFile("X509Crl", readCrlFromFile) } test("File isn't a crl") { - PemFileTestUtils.testEmptyFile[InvalidPemFormatException, X509CRL]( + PemBytesTestUtils.testEmptyFile[InvalidPemFormatException, X509CRL]( "X509Crl", readCrlFromFile ) } test("File is garbage") { - val handler = PemFileTestUtils.newHandler() // Lines were manually deleted from a real crl file val tempFile = TempFile.fromResourcePath("/crl/csl-intermediate-garbage.crl") // deleteOnExit is handled by TempFile - + val crlFile = new X509CrlFile(tempFile) val tryCrl = crlFile.readX509Crl() - assertLogMessage(handler.get, tempFile.getName, "Incomplete BER/DER data.") assertCrlException(tryCrl) } test("File is a Certificate Revocation List") { val tempFile = TempFile.fromResourcePath("/crl/csl-intermediate.crl") // deleteOnExit is handled by TempFile - + val crlFile = new X509CrlFile(tempFile) val tryCrl = crlFile.readX509Crl() diff --git a/util-security/src/test/scala/com/twitter/util/security/X509TrustManagerFactoryTest.scala b/util-security/src/test/scala/com/twitter/util/security/X509TrustManagerFactoryTest.scala index 3642c18361..3e5b4d09fb 100644 --- a/util-security/src/test/scala/com/twitter/util/security/X509TrustManagerFactoryTest.scala +++ b/util-security/src/test/scala/com/twitter/util/security/X509TrustManagerFactoryTest.scala @@ -4,12 +4,9 @@ import com.twitter.io.TempFile import java.io.File import java.security.cert.X509Certificate import javax.net.ssl.{TrustManager, X509TrustManager} -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class X509TrustManagerFactoryTest extends FunSuite { +class X509TrustManagerFactoryTest extends AnyFunSuite { private[this] def loadCertFromResource(resourcePath: String): X509Certificate = { val tempFile = TempFile.fromResourcePath(resourcePath) diff --git a/util-slf4j-api/BUILD b/util-slf4j-api/BUILD index b123da30a3..f51e55cc98 100644 --- a/util-slf4j-api/BUILD +++ b/util-slf4j-api/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-slf4j-api/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-slf4j-api/src/test/scala", ], diff --git a/util-slf4j-api/PROJECT b/util-slf4j-api/PROJECT new file mode 100644 index 0000000000..563cf7daa4 --- /dev/null +++ b/util-slf4j-api/PROJECT @@ -0,0 +1,8 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan + - mdickinson +watchers: + - csl-team@twitter.com diff --git a/util-slf4j-api/README.md b/util-slf4j-api/README.md index 69c986d73b..9e1b612e9a 100644 --- a/util-slf4j-api/README.md +++ b/util-slf4j-api/README.md @@ -1,43 +1,70 @@ # Twitter Util SLF4J API Support -This library provides a Scala-friendly logging wrapper around the [slf4j-api](https://www.slf4j.org/) [Logger](https://www.slf4j.org/apidocs/org/slf4j/Logger.html) interface. +This library provides a Scala-friendly logging wrapper around the [slf4j-api](https://www.slf4j.org/) +[Logger](https://www.slf4j.org/apidocs/org/slf4j/Logger.html) interface. ## Differences with `util-logging` -The code in `util-logging` is a scala wrapper over `java.util.logging` (JUL) which is a logging API **and** implementation. `util-slf4j-api` is an API-only and adds a thin scala wrapper over the [slf4j-api](http://www.slf4j.org/) which allows -users to choose their logging implementation. It is recommended that users move to using `util-slf4j-api` over `util-logging`. +The code in `util-logging` is a scala wrapper over `java.util.logging` (JUL) which is a logging API +**and** implementation. `util-slf4j-api` is an API-only and adds a thin scala wrapper over the +[slf4j-api](https://www.slf4j.org/) which allows users to choose their logging implementation. It is +recommended that users move to using `util-slf4j-api` over `util-logging`. -Since the [slf4j-api](https://www.slf4j.org/manual.html) is only an interface it requires an actual logging implementation. You should ensure that you do not end-up with *multiple* logging implementations on your classpath, e.g., you should not have multiple SLF4J bindings (`slf4j-nop`, `slf4j-log4j12`, `slf4j-jdk14`, etc.) and/or a `java.util.logging` implementation, etc. on your classpath as these are all competing implementations and since classpath order is non-deterministic, can lead to unexpected logging behavior. +Since the [slf4j-api](https://www.slf4j.org/manual.html) is only an interface it requires an actual +logging implementation. You should ensure that you do not end-up with *multiple* logging +implementations on your classpath, e.g., you should not have multiple SLF4J bindings (`slf4j-nop`, +`slf4j-log4j12`, `slf4j-jdk14`, etc.) and/or a `java.util.logging` implementation, etc. on your +classpath as these are all competing implementations and since classpath order is non-deterministic, +can lead to unexpected logging behavior. ## Usages -You can create a `c.t.util.logging.Logger` in different ways. You can pass a `name` to the `Logger#apply` factory method on the `c.t.util.logging.Logger` companion object: +You can create a `c.t.util.logging.Logger` in different ways. You can pass a `name` to the +`Logger#apply` factory method on the `c.t.util.logging.Logger` companion object: ```scala -val logger = Logger("name") +val logger = Logger("name") // scala +``` + +or in Java, use the `Logger#getLogger` method: + +```java +Logger log = Logger.getLogger("name"); // java ``` Or, pass in a class: ```scala -val logger = Logger(classOf[MyClass]) +val logger = Logger(classOf[MyClass]) // scala +``` + +```java +Logger log = Logger.getLogger(MyClass.class); // java ``` Or, use the runtime class wrapped by the implicit ClassTag parameter: ```scala -val logger = Logger[MyClass] +val logger = Logger[MyClass] // scala ``` Or, finally, use an SLF4J Logger instance: ```scala -val logger = Logger(LoggerFactory.getLogger("name")) +val logger = Logger(LoggerFactory.getLogger("name")) // scala +``` + +```java +Logger log = Logger.getLogger(LoggerFactory.getLogger("name")); // java ``` -There is also the `c.t.util.logging.Logging` trait which can be mixed into a class. Note, that the underlying SLF4J logger will be named according to the leaf class of any hierarchy into which the trait is mixed: +There is also the `c.t.util.logging.Logging` trait which can be mixed into a class. Note, that the +underlying SLF4J logger will be named according to the leaf class of any hierarchy into which the +trait is mixed: ```scala +package com.twitter + class Stock(symbol: String, price: BigDecimal) extends Logging { info(f"New stock with symbol = $symbol and price = $price%1.2f") ... @@ -47,11 +74,13 @@ class Stock(symbol: String, price: BigDecimal) extends Logging { This would produce a logging statement along the lines of: ``` -2017-01-05 13:05:49,349 INF Stock New stock with symbol = ASDF and price = 100.00 +2017-01-05 13:05:49,349 INF c.t.Stock New stock with symbol = ASDF and price = 100.00 ``` -If you are using Java and happen to extend a Scala class which mixes in the `c.t.util.logging.Logging` trait, you can either access the logger directly or choose to instantiate a new logger in your -Java class to log your statements. Again, note that the underlying SLF4J logger when mixed in via the trait will be named according to the leaf class into which it is mixed. +If you are using Java and happen to extend a Scala class which mixes in the `c.t.util.logging.Logging` +trait, you can either access the logger directly or choose to instantiate a new logger in your +Java class to log your statements. Again, note that the underlying SLF4J logger when mixed in via +the trait will be named according to the leaf class into which it is mixed. E.g., for a Scala class: @@ -61,7 +90,9 @@ class MyBaseScalaClass extends Logging { } ``` -You could choose to use the inherited logger directly since the inherited `c.t.util.logging.Logging` trait methods expect an anonymous function: +You could choose to use the inherited `logger` instance directly (since the inherited +`c.t.util.logging.Logging` trait methods would require using a `scala.Function0` as the trait +methods are explicitly call-by-name): ``` class MyJavaClass extends MyBaseScalaClass { @@ -73,20 +104,21 @@ class MyJavaClass extends MyBaseScalaClass { } ``` -Or simply instantiate a new logger: +Or instantiate a new logger explicitly: ``` class MyJavaClass extends MyBaseScalaClass { - private static final Logger logger = new Logger(this.getClass); + private static final Logger LOG = Logger.getLogger(MyJavaClass.class); ... public void methodThatDoesSomeLogging() { - logger.info("Accessing the superclass logger"); + LOG.info("Accessing the superclass logger"); } } ``` -For another Scala class that extends the above example Scala class with the mixed-in `c.t.util.logging.Logging` trait, you have the same options along with: +For another Scala class which extends the above example Scala class with the mixed-in +`c.t.util.logging.Logging` trait, you have the same options along with: - simply use the inherited Logging trait methods. - extend the `c.t.util.logging.Logging` trait again. diff --git a/util-slf4j-api/src/main/scala/BUILD b/util-slf4j-api/src/main/scala/BUILD index f22a32ae04..dc234d49d1 100644 --- a/util-slf4j-api/src/main/scala/BUILD +++ b/util-slf4j-api/src/main/scala/BUILD @@ -1,15 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-slf4j", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/slf4j:slf4j-api", - ], - exports = [ - "3rdparty/jvm/org/slf4j:slf4j-api", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", ], ) diff --git a/util-slf4j-api/src/main/scala/com/twitter/util/logging/BUILD b/util-slf4j-api/src/main/scala/com/twitter/util/logging/BUILD new file mode 100644 index 0000000000..7b6dfa2e56 --- /dev/null +++ b/util-slf4j-api/src/main/scala/com/twitter/util/logging/BUILD @@ -0,0 +1,17 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-slf4j", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/slf4j:slf4j-api", + ], + exports = [ + "3rdparty/jvm/org/slf4j:slf4j-api", + ], +) diff --git a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logger.scala b/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logger.scala index f78ed8c8c4..05d63455e0 100644 --- a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logger.scala +++ b/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logger.scala @@ -1,20 +1,39 @@ package com.twitter.util.logging import org.slf4j -import org.slf4j.{LoggerFactory, Marker} -import scala.reflect.{ClassTag, classTag} +import org.slf4j.LoggerFactory +import org.slf4j.Marker +import scala.reflect.ClassTag +import scala.reflect.classTag /** * Companion object for [[com.twitter.util.logging.Logger]] which provides * factory methods for creation. * - * @note Java users, see [[com.twitter.util.logging.Loggers.getLogger]] + * @define ApplyScaladocLink + * [[com.twitter.util.logging.Logger.apply(name:String):com\.twitter\.util\.logging\.Logger* apply(String)]] + * + * @define ApplyLoggerScaladocLink + * [[com.twitter.util.logging.Logger.apply(underlying:org\.slf4j\.Logger):com\.twitter\.util\.logging\.Logger* apply(Logger)]] + * + * @define ApplyClassScaladocLink + * [[com.twitter.util.logging.Logger.apply(clazz:Class[_]):com\.twitter\.util\.logging\.Logger* apply(Class)]] + * + * @define GetLoggerScaladocLink + * [[com.twitter.util.logging.Logger.getLogger(name:String):com\.twitter\.util\.logging\.Logger* getLogger(String)]] + * + * @define GetLoggerClassScaldocLink + * [[com.twitter.util.logging.Logger.getLogger(clazz:Class[_]):com\.twitter\.util\.logging\.Logger* getLogger(Class)]] + * + * @define GetLoggerLoggerScaladocLink + * [[com.twitter.util.logging.Logger.getLogger(underlying:org\.slf4j\.Logger):com\.twitter\.util\.logging\.Logger* getLogger(Logger)]] */ object Logger { /** * Create a [[com.twitter.util.logging.Logger]] for the given name. - * @param name name of the underlying Logger. + * @param name name of the underlying `org.slf4j.Logger`. + * @see Java users see $GetLoggerScaladocLink * * {{{ * val logger = Logger("name") @@ -24,35 +43,70 @@ object Logger { new Logger(LoggerFactory.getLogger(name)) } + /** + * Create a [[Logger]] for the given name. + * @param name name of the underlying `org.slf4j.Logger`. + * @see Scala users see $ApplyScaladocLink + * + * {{{ + * Logger logger = Logger.getLogger("name"); + * }}} + */ + def getLogger(name: String): Logger = + Logger(name) + /** * Create a [[com.twitter.util.logging.Logger]] named for the given class. * @param clazz class to use for naming the underlying Logger. + * @note Scala singleton object classes are converted to loggers named without the `$` suffix. + * @see Java users see $GetLoggerClassScaladocLink + * @see [[https://docs.scala-lang.org/tour/singleton-objects.html]] * * {{{ * val logger = Logger(classOf[MyClass]) * }}} */ def apply(clazz: Class[_]): Logger = { - new Logger(LoggerFactory.getLogger(clazz)) + if (isSingletonObject(clazz)) + new Logger(LoggerFactory.getLogger(objectClazzName(clazz))) + else new Logger(LoggerFactory.getLogger(clazz)) } + /** + * Create a [[Logger]] named for the given class. + * @param clazz class to use for naming the underlying `org.slf4j.Logger`. + * @see Scala users see $ApplyClassScaladocLink + * + * {{{ + * Logger logger = Logger.getLogger(MyClass.class); + * }}} + */ + def getLogger(clazz: Class[_]): Logger = + Logger(clazz) + /** * Create a [[com.twitter.util.logging.Logger]] for the runtime class wrapped - * by the implicit [[scala.reflect.ClassTag]] denoting the runtime class `T`. + * by the implicit `scala.reflect.ClassTag` denoting the runtime class `T`. * @tparam T the runtime class type + * @note Scala singleton object classes are converted to loggers named without the `$` suffix. + * @see [[https://docs.scala-lang.org/tour/singleton-objects.html]] * * {{{ * val logger = Logger[MyClass] * }}} */ def apply[T: ClassTag]: Logger = { - new Logger(LoggerFactory.getLogger(classTag[T].runtimeClass)) + val clazz = classTag[T].runtimeClass + if (isSingletonObject(clazz)) + new Logger(LoggerFactory.getLogger(objectClazzName(clazz))) + else new Logger(LoggerFactory.getLogger(clazz)) } /** * Create a [[com.twitter.util.logging.Logger]] wrapping the given underlying - * [[org.slf4j.Logger]]. + * `org.slf4j.Logger`. * @param underlying an `org.slf4j.Logger` + * @see Java users see $GetLoggerLoggerScaladocLink * * {{{ * val logger = Logger(LoggerFactory.getLogger("name")) @@ -61,14 +115,57 @@ object Logger { def apply(underlying: slf4j.Logger): Logger = { new Logger(underlying) } + + /** + * Create a [[Logger]] wrapping the given underlying + * `org.slf4j.Logger`. + * @param underlying an `org.slf4j.Logger` + * @see Scala users see $ApplyLoggerScaladocLink + * + * {{{ + * Logger logger = Logger.getLogger(LoggerFactory.getLogger("name")); + * }}} + */ + def getLogger(underlying: slf4j.Logger): Logger = + Logger(underlying) + + /** + * Returns true if the given instance is a Scala singleton object, false otherwise. + * EXPOSED FOR TESTING + * @see [[https://docs.scala-lang.org/tour/singleton-objects.html]] + */ + private[logging] def isSingletonObject(clazz: Class[_]): Boolean = { + // While Scala's reflection utilities offer this in a succinct and convenient + // method, the startup costs are quite high and in testing, the results not always + // accurate. The approach used here, while fragile to Scala internal changes, was + // deemed worth the tradeoff. This assumes that singleton objects have class names that + // end with $ and have a field on them called MODULE$. For example `object Foo` + // has the class name `Foo` and a `MODULE$` field. + // This works as of Scala 2.12, 2.13 and 3.0.2-RC1. + val className = clazz.getName + if (!className.endsWith("$")) + return false + + try { + clazz.getField("MODULE$") // this throws if the field doesn't exist + true + } catch { + case _: NoSuchFieldException => false + } + } + + // EXPOSED FOR TESTING + private[logging] def objectClazzName(clazz: Class[_]) = clazz.getName.stripSuffix("$") } /** - * A scala wrapper over a [[org.slf4j.Logger]]. + * A scala wrapper over a `org.slf4j.Logger`. + * + * The Logger extends [[https://docs.oracle.com/javase/8/docs/api/java/io/Serializable.html java.io.Serializable]] + * to support it's usage through the [[com.twitter.util.logging.Logging]] trait when the trait is mixed + * into a [[https://docs.oracle.com/javase/8/docs/api/java/io/Serializable.html java.io.Serializable]] class. * - * The Logger is [[Serializable]] to support it's usage through the - * [[com.twitter.util.logging.Logging]] trait when the trait is mixed - * into a [[Serializable]] class. + * Note, logging methods have a call-by-name variation which are intended for use from Scala. * * @define isLevelEnabled * @@ -76,8 +173,8 @@ object Logger { * * @define isLevelEnabledMarker * - * Determines if the named log level is enabled taking into consideration the given [[Marker]] data. - * Returns `true` if enabled, `false` otherwise. + * Determines if the named log level is enabled taking into consideration the given + * `org.slf4j.Marker` data. Returns `true` if enabled, `false` otherwise. * * @define log * @@ -86,18 +183,29 @@ object Logger { * @define logMarker * * Logs the given message at the named log level taking into consideration the - * given [[Marker]] data. + * given `org.slf4j.Marker` data. * * @define logWith * * Log the given parameterized message at the named log level using the given - * args. See [[http://www.slf4j.org/faq.html#logging_performance Parameterized Message Logging]] + * args. See [[https://www.slf4j.org/faq.html#logging_performance Parameterized Message Logging]] * * @define logWithMarker * * Log the given parameterized message at the named log level using the given - * args taking into consideration the given [[Marker]] data. - * See [[http://www.slf4j.org/faq.html#logging_performance Parameterized Message Logging]] + * args taking into consideration the given `org.slf4j.Marker` data. See [[https://www.slf4j.org/faq.html#logging_performance Parameterized Message Logging]] + * + * @define logResult + * + * Log the given message at the named log level formatted with the result of the + * passed in function using the underlying logger. The incoming string message should + * contain a single `%s` which will be replaced with the result[T] of the given function. + * + * Example: + * + * {{{ + * infoResult("The answer is: %s") {"42"} + * }}} */ @SerialVersionUID(1L) final class Logger private (underlying: slf4j.Logger) extends Serializable { @@ -125,18 +233,34 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { def trace(message: String): Unit = if (underlying.isTraceEnabled) underlying.trace(message) + /** $log */ + def trace(message: => Any): Unit = + if (underlying.isTraceEnabled) underlying.trace(toString(message)) + /** $log */ def trace(message: String, cause: Throwable): Unit = if (underlying.isTraceEnabled) underlying.trace(message, cause) + /** $log */ + def trace(message: => Any, cause: Throwable): Unit = + if (underlying.isTraceEnabled) underlying.trace(toString(message), cause) + /** $logMarker */ def trace(marker: Marker, message: String): Unit = if (underlying.isTraceEnabled(marker)) underlying.trace(marker, message) + /** $logMarker */ + def trace(marker: Marker, message: => Any): Unit = + if (underlying.isTraceEnabled(marker)) underlying.trace(marker, toString(message)) + /** $logMarker */ def trace(marker: Marker, message: String, cause: Throwable): Unit = if (underlying.isTraceEnabled(marker)) underlying.trace(marker, message, cause) + /** $logMarker */ + def trace(marker: Marker, message: => Any, cause: Throwable): Unit = + if (underlying.isTraceEnabled(marker)) underlying.trace(marker, toString(message), cause) + /** $logWith */ def traceWith(message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -145,6 +269,14 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isTraceEnabled) underlying.trace(message, args: _*) } + /** $logWith */ + def traceWith(message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isTraceEnabled) underlying.trace(toString(message)) + } else { + if (underlying.isTraceEnabled) underlying.trace(toString(message), args: _*) + } + /** $logWithMarker */ def traceWith(marker: Marker, message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -153,6 +285,21 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isTraceEnabled(marker)) underlying.trace(marker, message, args: _*) } + /** $logWithMarker */ + def traceWith(marker: Marker, message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isTraceEnabled(marker)) underlying.trace(marker, toString(message)) + } else { + if (underlying.isTraceEnabled(marker)) underlying.trace(marker, toString(message), args: _*) + } + + /** $logResult */ + def traceResult[T](message: => String)(fn: => T): T = { + val result = fn + trace(message.format(result)) + result + } + /* Debug */ /** $isLevelEnabled */ @@ -166,18 +313,34 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { def debug(message: String): Unit = if (underlying.isDebugEnabled) underlying.debug(message) + /** $log */ + def debug(message: => Any): Unit = + if (underlying.isDebugEnabled) underlying.debug(toString(message)) + /** $log */ def debug(message: String, cause: Throwable): Unit = if (underlying.isDebugEnabled) underlying.debug(message, cause) + /** $log */ + def debug(message: => Any, cause: Throwable): Unit = + if (underlying.isDebugEnabled) underlying.debug(toString(message), cause) + /** $logMarker */ def debug(marker: Marker, message: String): Unit = if (underlying.isDebugEnabled(marker)) underlying.debug(marker, message) + /** $logMarker */ + def debug(marker: Marker, message: => Any): Unit = + if (underlying.isDebugEnabled(marker)) underlying.debug(marker, toString(message)) + /** $logMarker */ def debug(marker: Marker, message: String, cause: Throwable): Unit = if (underlying.isDebugEnabled(marker)) underlying.debug(marker, message, cause) + /** $logMarker */ + def debug(marker: Marker, message: => Any, cause: Throwable): Unit = + if (underlying.isDebugEnabled(marker)) underlying.debug(marker, toString(message), cause) + /** $logWith */ def debugWith(message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -186,6 +349,14 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isDebugEnabled) underlying.debug(message, args: _*) } + /** $logWith */ + def debugWith(message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isDebugEnabled) underlying.debug(toString(message)) + } else { + if (underlying.isDebugEnabled) underlying.debug(toString(message), args: _*) + } + /** $logWithMarker */ def debugWith(marker: Marker, message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -194,6 +365,21 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isDebugEnabled(marker)) underlying.debug(marker, message, args: _*) } + /** $logWithMarker */ + def debugWith(marker: Marker, message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isDebugEnabled(marker)) underlying.debug(marker, toString(message)) + } else { + if (underlying.isDebugEnabled(marker)) underlying.debug(marker, toString(message), args: _*) + } + + /** $logResult */ + def debugResult[T](message: => String)(fn: => T): T = { + val result = fn + debug(message.format(result)) + result + } + /* Info */ /** $isLevelEnabled */ @@ -207,18 +393,34 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { def info(message: String): Unit = if (underlying.isInfoEnabled) underlying.info(message) + /** $log */ + def info(message: => Any): Unit = + if (underlying.isInfoEnabled) underlying.info(toString(message)) + /** $log */ def info(message: String, cause: Throwable): Unit = if (underlying.isInfoEnabled) underlying.info(message, cause) + /** $log */ + def info(message: => Any, cause: Throwable): Unit = + if (underlying.isInfoEnabled) underlying.info(toString(message), cause) + /** $logMarker */ def info(marker: Marker, message: String): Unit = if (underlying.isInfoEnabled(marker)) underlying.info(marker, message) + /** $logMarker */ + def info(marker: Marker, message: => Any): Unit = + if (underlying.isInfoEnabled(marker)) underlying.info(marker, toString(message)) + /** $logMarker */ def info(marker: Marker, message: String, cause: Throwable): Unit = if (underlying.isInfoEnabled(marker)) underlying.info(marker, message, cause) + /** $logMarker */ + def info(marker: Marker, message: => Any, cause: Throwable): Unit = + if (underlying.isInfoEnabled(marker)) underlying.info(marker, toString(message), cause) + /** $logWith */ def infoWith(message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -227,6 +429,14 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isInfoEnabled) underlying.info(message, args: _*) } + /** $logWith */ + def infoWith(message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isInfoEnabled) underlying.info(toString(message)) + } else { + if (underlying.isInfoEnabled) underlying.info(toString(message), args: _*) + } + /** $logWithMarker */ def infoWith(marker: Marker, message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -235,6 +445,21 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isInfoEnabled(marker)) underlying.info(marker, message, args: _*) } + /** $logWithMarker */ + def infoWith(marker: Marker, message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isInfoEnabled(marker)) underlying.info(marker, toString(message)) + } else { + if (underlying.isInfoEnabled(marker)) underlying.info(marker, toString(message), args: _*) + } + + /** $logResult */ + def infoResult[T](message: => String)(fn: => T): T = { + val result = fn + info(message.format(result)) + result + } + /* Warn */ /** $isLevelEnabled */ @@ -248,18 +473,34 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { def warn(message: String): Unit = if (underlying.isWarnEnabled) underlying.warn(message) + /** $log */ + def warn(message: => Any): Unit = + if (underlying.isWarnEnabled) underlying.warn(toString(message)) + /** $log */ def warn(message: String, cause: Throwable): Unit = if (underlying.isWarnEnabled) underlying.warn(message, cause) + /** $log */ + def warn(message: => Any, cause: Throwable): Unit = + if (underlying.isWarnEnabled) underlying.warn(toString(message), cause) + /** $logMarker */ def warn(marker: Marker, message: String): Unit = if (underlying.isWarnEnabled(marker)) underlying.warn(marker, message) + /** $logMarker */ + def warn(marker: Marker, message: => Any): Unit = + if (underlying.isWarnEnabled(marker)) underlying.warn(marker, toString(message)) + /** $logMarker */ def warn(marker: Marker, message: String, cause: Throwable): Unit = if (underlying.isWarnEnabled(marker)) underlying.warn(marker, message, cause) + /** $logMarker */ + def warn(marker: Marker, message: => Any, cause: Throwable): Unit = + if (underlying.isWarnEnabled(marker)) underlying.warn(marker, toString(message), cause) + /** $logWith */ def warnWith(message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -268,6 +509,14 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isWarnEnabled) underlying.warn(message, args: _*) } + /** $logWith */ + def warnWith(message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isWarnEnabled) underlying.warn(toString(message)) + } else { + if (underlying.isWarnEnabled) underlying.warn(toString(message), args: _*) + } + /** $logWithMarker */ def warnWith(marker: Marker, message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -276,6 +525,21 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isWarnEnabled(marker)) underlying.warn(marker, message, args: _*) } + /** $logWithMarker */ + def warnWith(marker: Marker, message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isWarnEnabled(marker)) underlying.warn(marker, toString(message)) + } else { + if (underlying.isWarnEnabled(marker)) underlying.warn(marker, toString(message), args: _*) + } + + /** $logResult */ + def warnResult[T](message: => String)(fn: => T): T = { + val result = fn + warn(message.format(result)) + result + } + /* Error */ /** $isLevelEnabled */ @@ -289,18 +553,34 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { def error(message: String): Unit = if (underlying.isErrorEnabled) underlying.error(message) + /** $log */ + def error(message: => Any): Unit = + if (underlying.isErrorEnabled) underlying.error(toString(message)) + /** $log */ def error(message: String, cause: Throwable): Unit = if (underlying.isErrorEnabled) underlying.error(message, cause) + /** $log */ + def error(message: => Any, cause: Throwable): Unit = + if (underlying.isErrorEnabled) underlying.error(toString(message), cause) + /** $logMarker */ def error(marker: Marker, message: String): Unit = if (underlying.isErrorEnabled(marker)) underlying.error(marker, message) + /** $logMarker */ + def error(marker: Marker, message: => Any): Unit = + if (underlying.isErrorEnabled(marker)) underlying.error(marker, toString(message)) + /** $logMarker */ def error(marker: Marker, message: String, cause: Throwable): Unit = if (underlying.isErrorEnabled(marker)) underlying.error(marker, message, cause) + /** $logMarker */ + def error(marker: Marker, message: => Any, cause: Throwable): Unit = + if (underlying.isErrorEnabled(marker)) underlying.error(marker, toString(message), cause) + /** $logWith */ def errorWith(message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -309,6 +589,14 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { if (underlying.isErrorEnabled) underlying.error(message, args: _*) } + /** $logWith */ + def errorWith(message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isErrorEnabled) underlying.error(toString(message)) + } else { + if (underlying.isErrorEnabled) underlying.error(toString(message), args: _*) + } + /** $logWithMarker */ def errorWith(marker: Marker, message: String, args: AnyRef*): Unit = if (args.isEmpty) { @@ -316,4 +604,29 @@ final class Logger private (underlying: slf4j.Logger) extends Serializable { } else { if (underlying.isErrorEnabled(marker)) underlying.error(marker, message, args: _*) } + + /** $logWithMarker */ + def errorWith(marker: Marker, message: => Any, args: AnyRef*): Unit = + if (args.isEmpty) { + if (underlying.isErrorEnabled(marker)) underlying.error(marker, toString(message)) + } else { + if (underlying.isErrorEnabled(marker)) underlying.error(marker, toString(message), args: _*) + } + + /** $logResult */ + def errorResult[T](message: => String)(fn: => T): T = { + val result = fn + error(message.format(result)) + result + } + + /* Private */ + + /** + * Convert [[Any]] type to [[String]] to safely deal with nulls in the above functions. + */ + private[this] def toString(obj: Any): String = obj match { + case null => "null" + case _ => obj.toString + } } diff --git a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Loggers.scala b/util-slf4j-api/src/main/scala/com/twitter/util/logging/Loggers.scala deleted file mode 100644 index 3600f01f09..0000000000 --- a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Loggers.scala +++ /dev/null @@ -1,48 +0,0 @@ -package com.twitter.util.logging - -import org.slf4j - -/** - * For Java usability. - * - * @note Scala users see [[com.twitter.util.logging.Logger]] - */ -object Loggers { - - /** - * Create a [[Logger]] for the given name. - * - * @param name name of the underlying `slf4j.Logger`. - * - * {{{ - * val logger = Logger("name") - * }}} - */ - def getLogger(name: String): Logger = - Logger(name) - - /** - * Create a [[Logger]] named for the given class. - * - * @param clazz class to use for naming the underlying slf4j.Logger. - * - * {{{ - * val logger = Logger(classOf[MyClass]) - * }}} - */ - def getLogger(clazz: Class[_]): Logger = - Logger(clazz) - - /** - * Create a [[Logger]] wrapping the given underlying - * [[org.slf4j.Logger]]. - * - * @param underlying an `org.slf4j.Logger` - * - * {{{ - * val logger = Logger(LoggerFactory.getLogger("name")) - * }}} - */ - def getLogger(underlying: slf4j.Logger): Logger = - Logger(underlying) -} diff --git a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logging.scala b/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logging.scala index c58bae04ad..2b344b2557 100644 --- a/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logging.scala +++ b/util-slf4j-api/src/main/scala/com/twitter/util/logging/Logging.scala @@ -1,6 +1,6 @@ package com.twitter.util.logging -import org.slf4j.{LoggerFactory, Marker} +import org.slf4j.Marker import scala.language.implicitConversions /** @@ -30,7 +30,7 @@ import scala.language.implicitConversions * @define isLevelEnabledMarker * * Determines if the named log level is enabled on the underlying logger taking into - * consideration the given [[Marker]] data. Returns `true` if enabled, `false` otherwise. + * consideration the given `org.slf4j.Marker` data. Returns `true` if enabled, `false` otherwise. * * @define log * @@ -40,7 +40,7 @@ import scala.language.implicitConversions * @define logMarker * * Logs the given message at the named log level with the underlying logger taking into - * consideration the given [[Marker]] data. + * consideration the given `org.slf4j.Marker` data. * * @define logResult * @@ -56,8 +56,7 @@ import scala.language.implicitConversions */ trait Logging { - private[this] final lazy val _logger: Logger = - Logger(LoggerFactory.getLogger(getClass)) + private[this] final lazy val _logger: Logger = Logger(getClass) /** * Return the underlying [[com.twitter.util.logging.Logger]] diff --git a/util-slf4j-api/src/test/java/BUILD b/util-slf4j-api/src/test/java/BUILD deleted file mode 100644 index bf67c899bd..0000000000 --- a/util-slf4j-api/src/test/java/BUILD +++ /dev/null @@ -1,11 +0,0 @@ -junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/slf4j:slf4j-api", - "3rdparty/jvm/org/slf4j:slf4j-simple", - "util/util-slf4j-api/src/main/scala", - "util/util-slf4j-api/src/test/scala", - ], -) diff --git a/util-slf4j-api/src/test/java/BUILD.bazel b/util-slf4j-api/src/test/java/BUILD.bazel new file mode 100644 index 0000000000..2c3035dbb1 --- /dev/null +++ b/util-slf4j-api/src/test/java/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-slf4j-api/src/test/java/com/twitter/util/logging", + ], +) diff --git a/util-slf4j-api/src/test/java/com/twitter/util/logging/BUILD.bazel b/util-slf4j-api/src/test/java/com/twitter/util/logging/BUILD.bazel new file mode 100644 index 0000000000..815cc1b44d --- /dev/null +++ b/util-slf4j-api/src/test/java/com/twitter/util/logging/BUILD.bazel @@ -0,0 +1,20 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/slf4j:slf4j-api", + scoped( + "3rdparty/jvm/org/slf4j:slf4j-simple", + scope = "runtime", + ), + "util/util-slf4j-api/src/main/scala", + ], +) diff --git a/util-slf4j-api/src/test/java/com/twitter/util/logging/LoggersCompilationTest.java b/util-slf4j-api/src/test/java/com/twitter/util/logging/LoggerCompilationTest.java similarity index 53% rename from util-slf4j-api/src/test/java/com/twitter/util/logging/LoggersCompilationTest.java rename to util-slf4j-api/src/test/java/com/twitter/util/logging/LoggerCompilationTest.java index 1bfd6dea07..23c078fb9d 100644 --- a/util-slf4j-api/src/test/java/com/twitter/util/logging/LoggersCompilationTest.java +++ b/util-slf4j-api/src/test/java/com/twitter/util/logging/LoggerCompilationTest.java @@ -3,19 +3,19 @@ import org.junit.Test; import org.slf4j.LoggerFactory; -import static org.junit.Assert.*; +import static org.junit.Assert.assertEquals; -public final class LoggersCompilationTest { +public class LoggerCompilationTest { @Test public void testApply1() { /* By slf4j Logger */ Logger logger = - Loggers.getLogger(LoggerFactory.getLogger(this.getClass())); + Logger.getLogger(LoggerFactory.getLogger(this.getClass())); assertEquals( - "com.twitter.util.logging.LoggersCompilationTest", + "com.twitter.util.logging.LoggerCompilationTest", logger.name()); } @@ -24,10 +24,10 @@ public void testApply2() { /* By Class */ Logger logger = - Loggers.getLogger(LoggersCompilationTest.class); + Logger.getLogger(LoggerCompilationTest.class); assertEquals( - "com.twitter.util.logging.LoggersCompilationTest", + "com.twitter.util.logging.LoggerCompilationTest", logger.name()); } @@ -36,10 +36,10 @@ public void testApply3() { /* By String */ Logger logger = - Loggers.getLogger("LoggersCompilationTest"); + Logger.getLogger("LoggerCompilationTest"); assertEquals( - "LoggersCompilationTest", + "LoggerCompilationTest", logger.name()); } } diff --git a/util-slf4j-api/src/test/resources/BUILD b/util-slf4j-api/src/test/resources/BUILD index c0cac6941d..99b7ffad66 100644 --- a/util-slf4j-api/src/test/resources/BUILD +++ b/util-slf4j-api/src/test/resources/BUILD @@ -1,3 +1,4 @@ resources( - sources = globs("*"), + sources = ["*"], + tags = ["bazel-compatible"], ) diff --git a/util-slf4j-api/src/test/scala/BUILD b/util-slf4j-api/src/test/scala/BUILD index 181d695241..266de50be1 100644 --- a/util-slf4j-api/src/test/scala/BUILD +++ b/util-slf4j-api/src/test/scala/BUILD @@ -1,13 +1,6 @@ -junit_tests( - sources = rglobs("*.scala", "*.java"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "3rdparty/jvm/org/slf4j:slf4j-api", - "3rdparty/jvm/org/slf4j:slf4j-simple", - "util/util-core/src/main/scala", - "util/util-slf4j-api/src/main/scala", - "util/util-slf4j-api/src/test/resources", + "util/util-slf4j-api/src/test/scala/com/twitter/util/logging", ], ) diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/BUILD b/util-slf4j-api/src/test/scala/com/twitter/util/logging/BUILD new file mode 100644 index 0000000000..72a6b2ff46 --- /dev/null +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/BUILD @@ -0,0 +1,50 @@ +scala_library( + name = "testutil", + sources = [ + "!*Test.scala", + "*.java", + "*.scala", + ], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "3rdparty/jvm/org/slf4j:slf4j-api", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-slf4j-api/src/test/resources", + ], + exports = [ + "3rdparty/jvm/org/slf4j:slf4j-api", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-slf4j-api/src/test/resources", + ], +) + +junit_tests( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + ":testutil", + "3rdparty/jvm/org/mockito:mockito-scala", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/src/jvm/org/scalatestplus/mockito", + scoped( + "3rdparty/jvm/org/slf4j:slf4j-simple", + scope = "runtime", + ), + "util/util-mock/src/main/scala/com/twitter/util/mock", + ], +) diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/Item.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/Item.scala index f1a5f91197..a4a8a41a28 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/Item.scala +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/Item.scala @@ -1,7 +1,5 @@ package com.twitter.util.logging -import com.twitter.util.logging.Item._ - object Item extends Logging { def baz: String = { @@ -11,6 +9,7 @@ object Item extends Logging { } class Item(val name: String, val description: String, size: Int) extends Logging { + import Item._ info(s"New item: name = $name, description = $description, size = $size") def dimension: Int = { diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaExtension.java b/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaExtension.java new file mode 100644 index 0000000000..60591e0cad --- /dev/null +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaExtension.java @@ -0,0 +1,25 @@ +package com.twitter.util.logging; + +import scala.runtime.AbstractFunction0; + +public class JavaExtension extends AbstractTraitWithLogging { + private static final Logger LOG = Logger.getLogger(JavaExtension.class); + + /** + * myMethod2 + * @return a "hello world" String + */ + public String myMethod2() { + info(new AbstractFunction0() { + @Override + public Object apply() { + return "In myMethod2: using trait#info"; + } + }); + + logger().info("In myMethod2: using logger#info"); + LOG.info("In myMethod2: using LOG#info"); + + return "Hello, world"; + } +} diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaItem.java b/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaItem.java index 6cdbf415b8..7abdc6c14d 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaItem.java +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/JavaItem.java @@ -1,7 +1,7 @@ package com.twitter.util.logging; class JavaItem extends SecondItem { - private static final Logger LOG = Logger.apply(JavaItem.class); + private static final Logger LOG = Logger.getLogger(JavaItem.class); public JavaItem(String name, String description, int size, int count) { super(name, description, size, count); diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/LoggingTest.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/LoggingTest.scala index 72481fde34..a5b3f4c511 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/LoggingTest.scala +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/LoggingTest.scala @@ -1,46 +1,46 @@ package com.twitter.util.logging -import org.mockito.Matchers._ import org.mockito.Mockito._ -import org.scalatest.mockito.MockitoSugar -import org.scalatest.{FunSuite, Matchers} -import org.slf4j -import scala.language.reflectiveCalls +import org.mockito.Mockito.when +import org.mockito.ArgumentMatchers.any +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.mockito.MockitoSugar -class LoggingTest extends FunSuite with Matchers with MockitoSugar { +class LoggingTest extends AnyFunSuite with Matchers with MockitoSugar { /* Trace */ test("Logging#trace enabled") { - val f = fixture(_.isTraceEnabled, isEnabled = true) + val f = Fixture(_.isTraceEnabled, isEnabled = true) f.logger.trace(f.message) verify(f.underlying).trace(f.message) } test("Logging#trace not enabled") { - val f = fixture(_.isTraceEnabled, isEnabled = false) + val f = Fixture(_.isTraceEnabled, isEnabled = false) f.logger.trace(f.message) - verify(f.underlying, never).trace(anyString) + verify(f.underlying, never()).trace(any[String]) } test("Logging#trace enabled with message and cause") { - val f = fixture(_.isTraceEnabled, isEnabled = true) + val f = Fixture(_.isTraceEnabled, isEnabled = true) f.logger.trace(f.message, f.cause) verify(f.underlying).trace(f.message, f.cause) } test("Logging#trace not enabled with message and cause") { - val f = fixture(_.isTraceEnabled, isEnabled = false) + val f = Fixture(_.isTraceEnabled, isEnabled = false) f.logger.trace(f.message, f.cause) - verify(f.underlying, never).trace(anyString, anyObject) + verify(f.underlying, never()).trace(any[String], any) } test("Logging#trace enabled with parameters") { - val f = fixture(_.isTraceEnabled, isEnabled = true) + val f = Fixture(_.isTraceEnabled, isEnabled = true) f.logger.traceWith(f.message, f.arg1) verify(f.underlying).trace(f.message, Seq(f.arg1): _*) @@ -51,48 +51,48 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#trace not enabled with parameters") { - val f = fixture(_.isTraceEnabled, isEnabled = false) + val f = Fixture(_.isTraceEnabled, isEnabled = false) f.logger.traceWith(f.message, f.arg1) - verify(f.underlying, never).trace(f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).trace(f.message, Seq(f.arg1): _*) f.logger.traceWith(f.message, f.arg1, f.arg2) - verify(f.underlying, never).trace(f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).trace(f.message, Seq(f.arg1, f.arg2): _*) f.logger.traceWith(f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).trace(f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).trace(f.message, f.arg1, f.arg2, f.arg3) } /* Debug */ test("Logging#debug enabled") { - val f = fixture(_.isDebugEnabled, isEnabled = true) + val f = Fixture(_.isDebugEnabled, isEnabled = true) f.logger.debug(f.message) verify(f.underlying).debug(f.message) } test("Logging#debug not enabled") { - val f = fixture(_.isDebugEnabled, isEnabled = false) + val f = Fixture(_.isDebugEnabled, isEnabled = false) f.logger.debug(f.message) - verify(f.underlying, never).debug(anyString) + verify(f.underlying, never()).debug(any[String]) } test("Logging#debug enabled with message and cause") { - val f = fixture(_.isDebugEnabled, isEnabled = true) + val f = Fixture(_.isDebugEnabled, isEnabled = true) f.logger.debug(f.message, f.cause) verify(f.underlying).debug(f.message, f.cause) } test("Logging#debug not enabled with message and cause") { - val f = fixture(_.isDebugEnabled, isEnabled = false) + val f = Fixture(_.isDebugEnabled, isEnabled = false) f.logger.debug(f.message, f.cause) - verify(f.underlying, never).debug(anyString, anyObject) + verify(f.underlying, never()).debug(any[String], any) } test("Logging#debug enabled with parameters") { - val f = fixture(_.isDebugEnabled, isEnabled = true) + val f = Fixture(_.isDebugEnabled, isEnabled = true) f.logger.debugWith(f.message, f.arg1) verify(f.underlying).debug(f.message, Seq(f.arg1): _*) @@ -103,48 +103,48 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#debug not enabled with parameters") { - val f = fixture(_.isDebugEnabled, isEnabled = false) + val f = Fixture(_.isDebugEnabled, isEnabled = false) f.logger.debugWith(f.message, f.arg1) - verify(f.underlying, never).debug(f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).debug(f.message, Seq(f.arg1): _*) f.logger.debugWith(f.message, f.arg1, f.arg2) - verify(f.underlying, never).debug(f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).debug(f.message, Seq(f.arg1, f.arg2): _*) f.logger.debugWith(f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).debug(f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).debug(f.message, f.arg1, f.arg2, f.arg3) } /* Info */ test("Logging#info enabled") { - val f = fixture(_.isInfoEnabled, isEnabled = true) + val f = Fixture(_.isInfoEnabled, isEnabled = true) f.logger.info(f.message) verify(f.underlying).info(f.message) } test("Logging#info not enabled") { - val f = fixture(_.isInfoEnabled, isEnabled = false) + val f = Fixture(_.isInfoEnabled, isEnabled = false) f.logger.info(f.message) - verify(f.underlying, never).info(anyString) + verify(f.underlying, never()).info(any[String]) } test("Logging#info enabled with message and cause") { - val f = fixture(_.isInfoEnabled, isEnabled = true) + val f = Fixture(_.isInfoEnabled, isEnabled = true) f.logger.info(f.message, f.cause) verify(f.underlying).info(f.message, f.cause) } test("Logging#info not enabled with message and cause") { - val f = fixture(_.isInfoEnabled, isEnabled = false) + val f = Fixture(_.isInfoEnabled, isEnabled = false) f.logger.info(f.message, f.cause) - verify(f.underlying, never).info(anyString, anyObject) + verify(f.underlying, never()).info(any[String], any) } test("Logging#info enabled with parameters") { - val f = fixture(_.isInfoEnabled, isEnabled = true) + val f = Fixture(_.isInfoEnabled, isEnabled = true) f.logger.infoWith(f.message, f.arg1) verify(f.underlying).info(f.message, Seq(f.arg1): _*) @@ -155,48 +155,48 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#info not enabled with parameters") { - val f = fixture(_.isInfoEnabled, isEnabled = false) + val f = Fixture(_.isInfoEnabled, isEnabled = false) f.logger.infoWith(f.message, f.arg1) - verify(f.underlying, never).info(f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).info(f.message, Seq(f.arg1): _*) f.logger.infoWith(f.message, f.arg1, f.arg2) - verify(f.underlying, never).info(f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).info(f.message, Seq(f.arg1, f.arg2): _*) f.logger.infoWith(f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).info(f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).info(f.message, f.arg1, f.arg2, f.arg3) } /* Warn */ test("Logging#warn enabled") { - val f = fixture(_.isWarnEnabled, isEnabled = true) + val f = Fixture(_.isWarnEnabled, isEnabled = true) f.logger.warn(f.message) verify(f.underlying).warn(f.message) } test("Logging#warn not enabled") { - val f = fixture(_.isWarnEnabled, isEnabled = false) + val f = Fixture(_.isWarnEnabled, isEnabled = false) f.logger.warn(f.message) - verify(f.underlying, never).warn(anyString) + verify(f.underlying, never()).warn(any[String]) } test("Logging#warn enabled with message and cause") { - val f = fixture(_.isWarnEnabled, isEnabled = true) + val f = Fixture(_.isWarnEnabled, isEnabled = true) f.logger.warn(f.message, f.cause) verify(f.underlying).warn(f.message, f.cause) } test("Logging#warn not enabled with message and cause") { - val f = fixture(_.isWarnEnabled, isEnabled = false) + val f = Fixture(_.isWarnEnabled, isEnabled = false) f.logger.warn(f.message, f.cause) - verify(f.underlying, never).warn(anyString, anyObject) + verify(f.underlying, never()).warn(any[String], any) } test("Logging#warn enabled with parameters") { - val f = fixture(_.isWarnEnabled, isEnabled = true) + val f = Fixture(_.isWarnEnabled, isEnabled = true) f.logger.warnWith(f.message, f.arg1) verify(f.underlying).warn(f.message, Seq(f.arg1): _*) @@ -207,48 +207,48 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#warn not enabled with parameters") { - val f = fixture(_.isWarnEnabled, isEnabled = false) + val f = Fixture(_.isWarnEnabled, isEnabled = false) f.logger.warnWith(f.message, f.arg1) - verify(f.underlying, never).warn(f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).warn(f.message, Seq(f.arg1): _*) f.logger.warnWith(f.message, f.arg1, f.arg2) - verify(f.underlying, never).warn(f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).warn(f.message, Seq(f.arg1, f.arg2): _*) f.logger.warnWith(f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).warn(f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).warn(f.message, f.arg1, f.arg2, f.arg3) } /* Error */ test("Logging#error enabled") { - val f = fixture(_.isErrorEnabled, isEnabled = true) + val f = Fixture(_.isErrorEnabled, isEnabled = true) f.logger.error(f.message) verify(f.underlying).error(f.message) } test("Logging#error not enabled") { - val f = fixture(_.isErrorEnabled, isEnabled = false) + val f = Fixture(_.isErrorEnabled, isEnabled = false) f.logger.error(f.message) - verify(f.underlying, never).error(anyString) + verify(f.underlying, never()).error(any[String]) } test("Logging#error enabled with message and cause") { - val f = fixture(_.isErrorEnabled, isEnabled = true) + val f = Fixture(_.isErrorEnabled, isEnabled = true) f.logger.error(f.message, f.cause) verify(f.underlying).error(f.message, f.cause) } test("Logging#error not enabled with message and cause") { - val f = fixture(_.isErrorEnabled, isEnabled = false) + val f = Fixture(_.isErrorEnabled, isEnabled = false) f.logger.error(f.message, f.cause) - verify(f.underlying, never).error(anyString, anyObject) + verify(f.underlying, never()).error(any[String], any) } test("Logging#error enabled with parameters") { - val f = fixture(_.isErrorEnabled, isEnabled = true) + val f = Fixture(_.isErrorEnabled, isEnabled = true) f.logger.errorWith(f.message, f.arg1) verify(f.underlying).error(f.message, Seq(f.arg1): _*) @@ -259,14 +259,14 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#error not enabled with parameters") { - val f = fixture(_.isErrorEnabled, isEnabled = false) + val f = Fixture(_.isErrorEnabled, isEnabled = false) f.logger.errorWith(f.message, f.arg1) - verify(f.underlying, never).error(f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).error(f.message, Seq(f.arg1): _*) f.logger.errorWith(f.message, f.arg1, f.arg2) - verify(f.underlying, never).error(f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).error(f.message, Seq(f.arg1, f.arg2): _*) f.logger.errorWith(f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).error(f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).error(f.message, f.arg1, f.arg2, f.arg3) } /* Instance Logging */ @@ -307,7 +307,16 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#Java class with new logger") { - new TestJavaClass() + new TestJavaClass // INFO com.twitter.util.logging.TestJavaClass - Creating new TestJavaClass instance. + + val t = new TraitWithLogging {} + t.myMethod1 // INFO com.twitter.util.logging.LoggingTest$$anon$1 - In myMethod1 + + val je = new JavaExtension + je.myMethod1 // INFO com.twitter.util.logging.JavaExtension - In myMethod1 + je.myMethod2() // INFO com.twitter.util.logging.JavaExtension - In myMethod2: using trait#info + // INFO com.twitter.util.logging.JavaExtension - In myMethod2: using logger#info + // INFO com.twitter.util.logging.JavaExtension - In myMethod2: using LOG#info } test("Logging#Java class that extends Scala class with Logging trait and new logger") { @@ -317,7 +326,7 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Logging#null values") { - val logger = Logger(slf4j.LoggerFactory.getLogger(this.getClass)) + val logger = Logger(org.slf4j.LoggerFactory.getLogger(this.getClass)) logger.info(null) val list = Nil logger.info(list.toString) @@ -348,17 +357,16 @@ class LoggingTest extends FunSuite with Matchers with MockitoSugar { } /* Private */ - - private def fixture(isEnabledFn: slf4j.Logger => Boolean, isEnabled: Boolean) = { - new { - val message = "msg" - val cause = new RuntimeException("TEST EXCEPTION") - val arg1 = "arg1" - val arg2 = new Integer(1) - val arg3 = "arg3" - val underlying = mock[slf4j.Logger] - when(isEnabledFn(underlying)).thenReturn(isEnabled) - val logger = Logger(underlying) - } + private case class Fixture( + isEnabledFn: org.slf4j.Logger => Boolean, + isEnabled: Boolean) { + val message = "msg" + val cause = new RuntimeException("TEST EXCEPTION") + val arg1 = "arg1" + val arg2: Integer = Integer.valueOf(1) + val arg3 = "arg3" + val underlying: org.slf4j.Logger = mock[org.slf4j.Logger] + when(isEnabledFn(underlying)).thenReturn(isEnabled) + val logger: Logger = Logger(underlying) } } diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/MarkerLoggingTest.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/MarkerLoggingTest.scala index def4806048..c0d42436fc 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/MarkerLoggingTest.scala +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/MarkerLoggingTest.scala @@ -1,45 +1,45 @@ package com.twitter.util.logging -import org.mockito.Matchers._ import org.mockito.Mockito._ -import org.scalatest.mockito.MockitoSugar -import org.scalatest.{FunSuite, Matchers} -import org.slf4j +import org.mockito.Mockito.when +import org.mockito.ArgumentMatchers.any +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.mockito.MockitoSugar import org.slf4j.Marker -import scala.language.reflectiveCalls -class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { +class MarkerLoggingTest extends AnyFunSuite with Matchers with MockitoSugar { test("Marker Logging#trace enabled") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = true) f.logger.trace(f.marker, f.message) verify(f.underlying).trace(f.marker, f.message) } test("Marker Logging#trace not enabled") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = false) f.logger.trace(f.marker, f.message) - verify(f.underlying, never).trace(anyObject[Marker], anyString) + verify(f.underlying, never()).trace(any[Marker], any[String]) } test("Marker Logging#trace enabled with message and cause") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = true) f.logger.trace(f.marker, f.message, f.cause) verify(f.underlying).trace(f.marker, f.message, f.cause) } test("Marker Logging#trace not enabled with message and cause") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = false) f.logger.trace(f.marker, f.message, f.cause) - verify(f.underlying, never).trace(anyObject[Marker], anyString, anyObject) + verify(f.underlying, never()).trace(any[Marker], any[String], any) } test("Marker Logging#trace enabled with parameters") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = true) f.logger.traceWith(f.marker, f.message, f.arg1) verify(f.underlying).trace(f.marker, f.message, Seq(f.arg1): _*) @@ -50,48 +50,48 @@ class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Marker Logging#trace not enabled with parameters") { - val f = fixture(_.isTraceEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isTraceEnabled(any[Marker]), isEnabled = false) f.logger.traceWith(f.marker, f.message, f.arg1) - verify(f.underlying, never).trace(f.marker, f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).trace(f.marker, f.message, Seq(f.arg1): _*) f.logger.traceWith(f.marker, f.message, f.arg1, f.arg2) - verify(f.underlying, never).trace(f.marker, f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).trace(f.marker, f.message, Seq(f.arg1, f.arg2): _*) f.logger.traceWith(f.marker, f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).trace(f.marker, f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).trace(f.marker, f.message, f.arg1, f.arg2, f.arg3) } /* Debug */ test("Marker Logging#debug enabled") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = true) f.logger.debug(f.marker, f.message) verify(f.underlying).debug(f.marker, f.message) } test("Marker Logging#debug not enabled") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = false) f.logger.debug(f.marker, f.message) - verify(f.underlying, never).debug(anyObject[Marker], anyString) + verify(f.underlying, never()).debug(any[Marker], any[String]) } test("Marker Logging#debug enabled with message and cause") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = true) f.logger.debug(f.marker, f.message, f.cause) verify(f.underlying).debug(f.marker, f.message, f.cause) } test("Marker Logging#debug not enabled with message and cause") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = false) f.logger.debug(f.marker, f.message, f.cause) - verify(f.underlying, never).debug(anyObject[Marker], anyString, anyObject) + verify(f.underlying, never()).debug(any[Marker], any[String], any) } test("Marker Logging#debug enabled with parameters") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = true) f.logger.debugWith(f.marker, f.message, f.arg1) verify(f.underlying).debug(f.marker, f.message, Seq(f.arg1): _*) @@ -102,48 +102,48 @@ class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Marker Logging#debug not enabled with parameters") { - val f = fixture(_.isDebugEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isDebugEnabled(any[Marker]), isEnabled = false) f.logger.debugWith(f.marker, f.message, f.arg1) - verify(f.underlying, never).debug(f.marker, f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).debug(f.marker, f.message, Seq(f.arg1): _*) f.logger.debugWith(f.marker, f.message, f.arg1, f.arg2) - verify(f.underlying, never).debug(f.marker, f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).debug(f.marker, f.message, Seq(f.arg1, f.arg2): _*) f.logger.debugWith(f.marker, f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).debug(f.marker, f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).debug(f.marker, f.message, f.arg1, f.arg2, f.arg3) } /* Info */ test("Marker Logging#info enabled") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = true) f.logger.info(f.marker, f.message) verify(f.underlying).info(f.marker, f.message) } test("Marker Logging#info not enabled") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = false) f.logger.info(f.marker, f.message) - verify(f.underlying, never).info(anyObject[Marker], anyString) + verify(f.underlying, never()).info(any[Marker], any[String]) } test("Marker Logging#info enabled with message and cause") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = true) f.logger.info(f.marker, f.message, f.cause) verify(f.underlying).info(f.marker, f.message, f.cause) } test("Marker Logging#info not enabled with message and cause") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = false) f.logger.info(f.marker, f.message, f.cause) - verify(f.underlying, never).info(anyObject[Marker], anyString, anyObject) + verify(f.underlying, never()).info(any[Marker], any[String], any) } test("Marker Logging#info enabled with parameters") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = true) f.logger.infoWith(f.marker, f.message, f.arg1) verify(f.underlying).info(f.marker, f.message, Seq(f.arg1): _*) @@ -154,48 +154,48 @@ class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Marker Logging#info not enabled with parameters") { - val f = fixture(_.isInfoEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isInfoEnabled(any[Marker]), isEnabled = false) f.logger.infoWith(f.marker, f.message, f.arg1) - verify(f.underlying, never).info(f.marker, f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).info(f.marker, f.message, Seq(f.arg1): _*) f.logger.infoWith(f.marker, f.message, f.arg1, f.arg2) - verify(f.underlying, never).info(f.marker, f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).info(f.marker, f.message, Seq(f.arg1, f.arg2): _*) f.logger.infoWith(f.marker, f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).info(f.marker, f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).info(f.marker, f.message, f.arg1, f.arg2, f.arg3) } /* Warn */ test("Marker Logging#warn enabled") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = true) f.logger.warn(f.marker, f.message) verify(f.underlying).warn(f.marker, f.message) } test("Marker Logging#warn not enabled") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = false) f.logger.warn(f.marker, f.message) - verify(f.underlying, never).warn(anyObject[Marker], anyString) + verify(f.underlying, never()).warn(any[Marker], any[String]) } test("Marker Logging#warn enabled with message and cause") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = true) f.logger.warn(f.marker, f.message, f.cause) verify(f.underlying).warn(f.marker, f.message, f.cause) } test("Marker Logging#warn not enabled with message and cause") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = false) f.logger.warn(f.marker, f.message, f.cause) - verify(f.underlying, never).warn(anyObject[Marker], anyString, anyObject) + verify(f.underlying, never()).warn(any[Marker], any[String], any) } test("Marker Logging#warn enabled with parameters") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = true) f.logger.warnWith(f.marker, f.message, f.arg1) verify(f.underlying).warn(f.marker, f.message, Seq(f.arg1): _*) @@ -206,48 +206,48 @@ class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Marker Logging#warn not enabled with parameters") { - val f = fixture(_.isWarnEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isWarnEnabled(any[Marker]), isEnabled = false) f.logger.warnWith(f.marker, f.message, f.arg1) - verify(f.underlying, never).warn(f.marker, f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).warn(f.marker, f.message, Seq(f.arg1): _*) f.logger.warnWith(f.marker, f.message, f.arg1, f.arg2) - verify(f.underlying, never).warn(f.marker, f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).warn(f.marker, f.message, Seq(f.arg1, f.arg2): _*) f.logger.warnWith(f.marker, f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).warn(f.marker, f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).warn(f.marker, f.message, f.arg1, f.arg2, f.arg3) } /* Error */ test("Marker Logging#error enabled") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = true) f.logger.error(f.marker, f.message) verify(f.underlying).error(f.marker, f.message) } test("Marker Logging#error not enabled") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = false) f.logger.error(f.marker, f.message) - verify(f.underlying, never).error(anyObject[Marker], anyString) + verify(f.underlying, never()).error(any[Marker], any[String]) } test("Marker Logging#error enabled with message and cause") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = true) f.logger.error(f.marker, f.message, f.cause) verify(f.underlying).error(f.marker, f.message, f.cause) } test("Marker Logging#error not enabled with message and cause") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = false) f.logger.error(f.marker, f.message, f.cause) - verify(f.underlying, never).error(anyObject[Marker], anyString, anyObject) + verify(f.underlying, never()).error(any[Marker], any[String], any) } test("Marker Logging#error enabled with parameters") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = true) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = true) f.logger.errorWith(f.marker, f.message, f.arg1) verify(f.underlying).error(f.marker, f.message, Seq(f.arg1): _*) @@ -258,30 +258,28 @@ class MarkerLoggingTest extends FunSuite with Matchers with MockitoSugar { } test("Marker Logging#error not enabled with parameters") { - val f = fixture(_.isErrorEnabled(anyObject[Marker]), isEnabled = false) + val f = Fixture(_.isErrorEnabled(any[Marker]), isEnabled = false) f.logger.errorWith(f.marker, f.message, f.arg1) - verify(f.underlying, never).error(f.marker, f.message, Seq(f.arg1): _*) + verify(f.underlying, never()).error(f.marker, f.message, Seq(f.arg1): _*) f.logger.errorWith(f.marker, f.message, f.arg1, f.arg2) - verify(f.underlying, never).error(f.marker, f.message, Seq(f.arg1, f.arg2): _*) + verify(f.underlying, never()).error(f.marker, f.message, Seq(f.arg1, f.arg2): _*) f.logger.errorWith(f.marker, f.message, f.arg1, f.arg2, f.arg3) - verify(f.underlying, never).error(f.marker, f.message, f.arg1, f.arg2, f.arg3) + verify(f.underlying, never()).error(f.marker, f.message, f.arg1, f.arg2, f.arg3) } /* Private */ - private def fixture(p: slf4j.Logger => Boolean, isEnabled: Boolean) = { - new { - val marker: Marker = TestMarker - val message = "msg" - val cause = new RuntimeException("TEST EXCEPTION") - val arg1 = "arg1" - val arg2 = new Integer(1) - val arg3 = "arg3" - val underlying: org.slf4j.Logger = mock[org.slf4j.Logger] - when(p(underlying)).thenReturn(isEnabled) - val logger = Logger(underlying) - } + private case class Fixture(p: org.slf4j.Logger => Boolean, isEnabled: Boolean) { + val marker: Marker = TestMarker + val message = "msg" + val cause = new RuntimeException("TEST EXCEPTION") + val arg1 = "arg1" + val arg2: Integer = Integer.valueOf(1) + val arg3 = "arg3" + val underlying: org.slf4j.Logger = mock[org.slf4j.Logger] + when(p(underlying)).thenReturn(isEnabled) + val logger: Logger = Logger(underlying) } object TestMarker extends Marker { diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/SerializableLoggingTest.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/SerializableLoggingTest.scala index ba77802a1e..2d06d59a67 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/SerializableLoggingTest.scala +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/SerializableLoggingTest.scala @@ -1,10 +1,10 @@ package com.twitter.util.logging -import com.twitter.util.TempFolder +import com.twitter.io.TempFolder import java.io._ -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class SerializableLoggingTest extends FunSuite with TempFolder { +class SerializableLoggingTest extends AnyFunSuite with TempFolder { def read[T](filename: String): T = { new ObjectInputStream(new FileInputStream(new File(folderName, filename))) diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/Stock.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/Stock.scala index e399716c52..8288e401e5 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/Stock.scala +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/Stock.scala @@ -1,6 +1,6 @@ package com.twitter.util.logging -class Stock(symbol: String, price: BigDecimal) extends Logging { +class Stock(symbol: String, price: Double) extends Logging { info(f"New stock with symbol = $symbol and price = $price%1.2f") def quote: String = f"$symbol = $price%1.2f" } diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/TestJavaClass.java b/util-slf4j-api/src/test/scala/com/twitter/util/logging/TestJavaClass.java index f3226abea2..31162aa430 100644 --- a/util-slf4j-api/src/test/scala/com/twitter/util/logging/TestJavaClass.java +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/TestJavaClass.java @@ -1,7 +1,7 @@ package com.twitter.util.logging; class TestJavaClass { - private static final Logger LOG = Logger.apply(TestJavaClass.class); + private static final Logger LOG = Logger.getLogger(TestJavaClass.class); public TestJavaClass() { LOG.info("Creating new TestJavaClass instance."); diff --git a/util-slf4j-api/src/test/scala/com/twitter/util/logging/TraitWithLogging.scala b/util-slf4j-api/src/test/scala/com/twitter/util/logging/TraitWithLogging.scala new file mode 100644 index 0000000000..e0553edaf1 --- /dev/null +++ b/util-slf4j-api/src/test/scala/com/twitter/util/logging/TraitWithLogging.scala @@ -0,0 +1,11 @@ +package com.twitter.util.logging + +abstract class AbstractTraitWithLogging extends TraitWithLogging + +trait TraitWithLogging extends Logging { + + def myMethod1: String = { + info("In myMethod1") + "Hello, World" + } +} diff --git a/util-slf4j-jul-bridge/BUILD b/util-slf4j-jul-bridge/BUILD index db92890fe8..41a89e4ae0 100644 --- a/util-slf4j-jul-bridge/BUILD +++ b/util-slf4j-jul-bridge/BUILD @@ -1,4 +1,5 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-slf4j-jul-bridge/src/main/scala", ], diff --git a/util-slf4j-jul-bridge/PROJECT b/util-slf4j-jul-bridge/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-slf4j-jul-bridge/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-slf4j-jul-bridge/src/main/scala/BUILD b/util-slf4j-jul-bridge/src/main/scala/BUILD index 03616c6c52..92db1a6284 100644 --- a/util-slf4j-jul-bridge/src/main/scala/BUILD +++ b/util-slf4j-jul-bridge/src/main/scala/BUILD @@ -1,15 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-slf4j-jul-bridge", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/slf4j:jul-to-slf4j", - "3rdparty/jvm/org/slf4j:slf4j-api", - "util/util-core/src/main/scala", - "util/util-slf4j-api/src/main/scala", + "util/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging", ], ) diff --git a/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/BUILD b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/BUILD new file mode 100644 index 0000000000..cad256b797 --- /dev/null +++ b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/BUILD @@ -0,0 +1,18 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-slf4j-jul-bridge", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/slf4j:jul-to-slf4j", + "3rdparty/jvm/org/slf4j:slf4j-api", + "util/util-app/src/main/scala", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + ], +) diff --git a/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridge.scala b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridge.scala new file mode 100644 index 0000000000..6fd6ec13b5 --- /dev/null +++ b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridge.scala @@ -0,0 +1,21 @@ +package com.twitter.util.logging + +import com.twitter.app.App + +/** + * Installing the `SLF4JBridgeHandler` should be done as early as possible in order to capture any + * log statements emitted to JUL. Generally it is best to install the handler in the constructor of + * the main class of your application. + * + * This trait can be mixed into a [[com.twitter.app.App]] type and will attempt to install the + * `SLF4JBridgeHandler` via the [[Slf4jBridgeUtility]] as part of the constructor of the resultant + * class. + * @see [[Slf4jBridgeUtility]] + * @see [[org.slf4j.bridge.SLF4JBridgeHandler]] + * @see [[https://www.slf4j.org/legacy.html Bridging legacy APIs]] + */ +trait Slf4jBridge { self: App => + + /** Attempt Slf4jBridgeHandler installation */ + Slf4jBridgeUtility.attemptSlf4jBridgeHandlerInstallation() +} diff --git a/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridgeUtility.scala b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridgeUtility.scala index f19b4df8be..a6168573ee 100644 --- a/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridgeUtility.scala +++ b/util-slf4j-jul-bridge/src/main/scala/com/twitter/util/logging/Slf4jBridgeUtility.scala @@ -5,32 +5,34 @@ import org.slf4j.LoggerFactory import org.slf4j.bridge.SLF4JBridgeHandler /** - * A utility to safely install the [[http://www.slf4j.org/ slf4j-api]] [[http://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]]. + * A utility to safely install the [[https://www.slf4j.org/ slf4j-api]] [[https://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]]. * The utility attempts to detect if the `slf4j-jdk14` dependency is present on * the classpath as it is unwise to install the `jul-to-slf4j` bridge with the `slf4j-jdk14` - * dependency present as it may cause an [[http://www.slf4j.org/legacy.html#julRecursion infinite loop]]. + * dependency present as it may cause an [[https://www.slf4j.org/legacy.html#julRecursion infinite loop]]. * - * If the [[http://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]] is already installed we do not try to install + * If the [[https://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]] is already installed we do not try to install * it again. * * Additionally note there is a performance impact for bridging jul log messages in this - * manner. However, using the [[http://logback.qos.ch/manual/configuration.html#LevelChangePropagator Logback LevelChangePropagator]] + * manner. However, using the [[https://logback.qos.ch/manual/configuration.html#LevelChangePropagator Logback LevelChangePropagator]] * (> 0.9.25) can eliminate this performance impact. * - * @see [[http://www.slf4j.org/legacy.html#jul-to-slf4j SLF4J: jul-to-slf4j bridge]] + * @see [[org.slf4j.bridge.SLF4JBridgeHandler]] + * @see [[https://www.slf4j.org/legacy.html Bridging legacy APIs]] + * @see [[https://www.slf4j.org/legacy.html#jul-to-slf4j SLF4J: jul-to-slf4j bridge]] */ object Slf4jBridgeUtility extends Logging { /** - * Attempt installation of the [[http://www.slf4j.org/ slf4j-api]] [[http://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]]. - * The [[http://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]] will not be installed if already installed or if the + * Attempt installation of the [[https://www.slf4j.org/ slf4j-api]] [[https://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]]. + * The [[https://www.slf4j.org/api/org/slf4j/bridge/SLF4JBridgeHandler.html SLF4JBridgeHandler]] will not be installed if already installed or if the * `slf4j-jdk14` dependency is detected on the classpath. */ def attemptSlf4jBridgeHandlerInstallation(): Unit = install() /* Private */ - private val install = Once { + private[this] val install = Once { if (!SLF4JBridgeHandler.isInstalled && canInstallBridgeHandler) { SLF4JBridgeHandler.removeHandlersForRootLogger() SLF4JBridgeHandler.install() @@ -42,17 +44,18 @@ object Slf4jBridgeUtility extends Logging { * We do not want to install the bridge handler if the * `org.slf4j.impl.JDK14LoggerFactory` exists on the classpath. */ - private def canInstallBridgeHandler: Boolean = { + private[this] def canInstallBridgeHandler: Boolean = { try { Class.forName("org.slf4j.impl.JDK14LoggerFactory", false, this.getClass.getClassLoader) LoggerFactory .getLogger(this.getClass) .warn( "Detected [org.slf4j.impl.JDK14LoggerFactory] on classpath. " + - "SLF4JBridgeHandler will not be installed, see: http://www.slf4j.org/legacy.html#julRecursion") + "SLF4JBridgeHandler will not be installed, see: https://www.slf4j.org/legacy.html#julRecursion" + ) false } catch { - case e: ClassNotFoundException => + case _: ClassNotFoundException => true } } diff --git a/util-stats/BUILD b/util-stats/BUILD index 0de6ae6acf..c5ae2b4be4 100644 --- a/util-stats/BUILD +++ b/util-stats/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-stats/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-stats/src/test/java", "util/util-stats/src/test/scala", diff --git a/util-stats/OWNERS b/util-stats/OWNERS deleted file mode 100644 index 37c64cba94..0000000000 --- a/util-stats/OWNERS +++ /dev/null @@ -1,6 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta -vkostyukov diff --git a/util-stats/PROJECT b/util-stats/PROJECT new file mode 100644 index 0000000000..fac8decd61 --- /dev/null +++ b/util-stats/PROJECT @@ -0,0 +1,6 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan + - observability-team:ldap diff --git a/util-stats/src/main/java/BUILD b/util-stats/src/main/java/BUILD new file mode 100644 index 0000000000..d046c97a6d --- /dev/null +++ b/util-stats/src/main/java/BUILD @@ -0,0 +1,6 @@ +target( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-stats/src/main/java/com/twitter/finagle/stats", + ], +) diff --git a/util-stats/src/main/java/com/twitter/finagle/stats/BUILD b/util-stats/src/main/java/com/twitter/finagle/stats/BUILD new file mode 100644 index 0000000000..245d3d6916 --- /dev/null +++ b/util-stats/src/main/java/com/twitter/finagle/stats/BUILD @@ -0,0 +1,14 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = artifact( + org = "com.twitter", + name = "util-stats-java", + repo = artifactory, + ), + tags = ["bazel-compatible"], + dependencies = [ + "util/util-stats", + ], +) diff --git a/util-stats/src/main/java/com/twitter/finagle/stats/MetricTypes.java b/util-stats/src/main/java/com/twitter/finagle/stats/MetricTypes.java new file mode 100644 index 0000000000..eed98aedb8 --- /dev/null +++ b/util-stats/src/main/java/com/twitter/finagle/stats/MetricTypes.java @@ -0,0 +1,20 @@ +package com.twitter.finagle.stats; + +/** + * Java APIs for {@link MetricBuilder.MetricType} + */ +public class MetricTypes { + private MetricTypes() { + throw new IllegalStateException(); + } + + public static final MetricBuilder.MetricType COUNTER_TYPE = MetricBuilder.CounterType$.MODULE$; + + public static final MetricBuilder.MetricType UNLATCHED_COUNTER_TYPE = MetricBuilder.UnlatchedCounter$.MODULE$; + + public static final MetricBuilder.MetricType COUNTERISH_GAUGE_TYPE = MetricBuilder.CounterishGaugeType$.MODULE$; + + public static final MetricBuilder.MetricType GAUGE_TYPE = MetricBuilder.GaugeType$.MODULE$; + + public static final MetricBuilder.MetricType HISTOGRAM_TYPE = MetricBuilder.HistogramType$.MODULE$; +} diff --git a/util-stats/src/main/scala/BUILD b/util-stats/src/main/scala/BUILD index c32122ce4a..fcd1b8ddc2 100644 --- a/util-stats/src/main/scala/BUILD +++ b/util-stats/src/main/scala/BUILD @@ -1,18 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-stats", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", - "3rdparty/jvm/com/google/code/findbugs:jsr305", - "util/util-core/src/main/scala", - "util/util-lint/src/main/scala", - ], - exports = [ - "3rdparty/jvm/com/github/ben-manes/caffeine", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", ], ) diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/BUILD b/util-stats/src/main/scala/com/twitter/finagle/stats/BUILD new file mode 100644 index 0000000000..9b31b3e727 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/BUILD @@ -0,0 +1,27 @@ +scala_library( + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-stats", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson:jackson-module-scala", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/datatype:jackson-datatype-joda", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/com/google/code/findbugs:jsr305", + "util/util-app/src/main/scala", + "util/util-core:util-core-util", + "util/util-lint/src/main/scala/com/twitter/util/lint", + ], + exports = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + ], +) diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/BlacklistStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/BlacklistStatsReceiver.scala deleted file mode 100644 index 0c1e4ed2a7..0000000000 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/BlacklistStatsReceiver.scala +++ /dev/null @@ -1,28 +0,0 @@ -package com.twitter.finagle.stats - -/** - * A blacklisting [[StatsReceiver]]. If the name for a metric is found to be - * blacklisted, nothing is recorded. - * - * @param self a base [[StatsReceiver]], used for metrics that aren't - * blacklisted - * @param blacklisted a predicate that reads a name and returns true to - * blacklist, and false to let it pass through - */ -class BlacklistStatsReceiver(protected val self: StatsReceiver, blacklisted: Seq[String] => Boolean) - extends StatsReceiverProxy { - - override def counter(verbosity: Verbosity, name: String*): Counter = - getStatsReceiver(name).counter(verbosity, name: _*) - - override def stat(verbosity: Verbosity, name: String*): Stat = - getStatsReceiver(name).stat(verbosity, name: _*) - - override def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge = - getStatsReceiver(name).addGauge(verbosity, name: _*)(f) - - private[this] def getStatsReceiver(name: Seq[String]): StatsReceiver = - if (blacklisted(name)) NullStatsReceiver else self - - override def toString: String = s"BlacklistStatsReceiver($self)" -} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala index ab3f0be28b..82c4b484e2 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/BroadcastStatsReceiver.scala @@ -1,5 +1,10 @@ package com.twitter.finagle.stats +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.util.Return +import com.twitter.util.Throw +import com.twitter.util.Try + /** * BroadcastStatsReceiver is a helper object that create a StatsReceiver wrapper around multiple * StatsReceivers (n). @@ -15,23 +20,41 @@ object BroadcastStatsReceiver { private class Two(first: StatsReceiver, second: StatsReceiver) extends StatsReceiver with DelegatingStatsReceiver { - val repr = this + val repr: AnyRef = this + + override def scopeTranslation: NameTranslatingStatsReceiver.Mode = + NameTranslatingStatsReceiver.FullTranslation - def counter(verbosity: Verbosity, names: String*): Counter = new BroadcastCounter.Two( - first.counter(verbosity, names: _*), - second.counter(verbosity, names: _*) + def counter(metricBuilder: MetricBuilder) = new BroadcastCounter.Two( + first.counter(metricBuilder), + second.counter(metricBuilder) ) - def stat(verbosity: Verbosity, names: String*): Stat = - new BroadcastStat.Two(first.stat(verbosity, names: _*), second.stat(verbosity, names: _*)) + def stat(metricBuilder: MetricBuilder) = new BroadcastStat.Two( + first.stat(metricBuilder), + second.stat(metricBuilder) + ) - def addGauge(verbosity: Verbosity, names: String*)(f: => Float): Gauge = new Gauge { - val firstGauge = first.addGauge(verbosity, names: _*)(f) - val secondGauge = second.addGauge(verbosity, names: _*)(f) + def addGauge(metricBuilder: MetricBuilder)(f: => Float) = new Gauge { + val firstGauge = first.addGauge(metricBuilder)(f) + val secondGauge = second.addGauge(metricBuilder)(f) def remove(): Unit = { firstGauge.remove() secondGauge.remove() } + def metadata: Metadata = MultiMetadata(Seq(firstGauge.metadata, secondGauge.metadata)) + } + + override def registerExpression( + expressionSchema: ExpressionSchema + ): Try[Unit] = { + Try.collect( + Seq( + first.registerExpression(expressionSchema), + second.registerExpression(expressionSchema))) match { + case Return(_) => Return.Unit + case Throw(throwable) => Throw(throwable) + } } def underlying: Seq[StatsReceiver] = Seq(first, second) @@ -41,19 +64,31 @@ object BroadcastStatsReceiver { } private class N(srs: Seq[StatsReceiver]) extends StatsReceiver with DelegatingStatsReceiver { - val repr = this + val repr: AnyRef = this + + override def scopeTranslation: NameTranslatingStatsReceiver.Mode = + NameTranslatingStatsReceiver.FullTranslation - def counter(verbosity: Verbosity, names: String*): Counter = - BroadcastCounter(srs.map { _.counter(verbosity, names: _*) }) + def counter(metricBuilder: MetricBuilder) = + BroadcastCounter(srs.map { _.counter(metricBuilder) }) - def stat(verbosity: Verbosity, names: String*): Stat = - BroadcastStat(srs.map { _.stat(verbosity, names: _*) }) + def stat(metricBuilder: MetricBuilder) = + BroadcastStat(srs.map { _.stat(metricBuilder) }) - def addGauge(verbosity: Verbosity, names: String*)(f: => Float): Gauge = new Gauge { - val gauges = srs.map { _.addGauge(verbosity, names: _*)(f) } + def addGauge(metricBuilder: MetricBuilder)(f: => Float) = new Gauge { + val gauges = srs.map { _.addGauge(metricBuilder)(f) } def remove(): Unit = gauges.foreach { _.remove() } + def metadata: Metadata = MultiMetadata(gauges.map(_.metadata)) } + override def registerExpression( + expressionSchema: ExpressionSchema + ): Try[Unit] = + Try.collect(srs.map(_.registerExpression(expressionSchema))) match { + case Return(_) => Return.Unit + case Throw(throwable) => Throw(throwable) + } + def underlying: Seq[StatsReceiver] = srs override def toString: String = @@ -78,6 +113,7 @@ object BroadcastCounter { private object NullCounter extends Counter { def incr(delta: Long): Unit = () + def metadata: Metadata = NoMetadata } private[stats] class Two(a: Counter, b: Counter) extends Counter { @@ -85,6 +121,7 @@ object BroadcastCounter { a.incr(delta) b.incr(delta) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata)) } private class Three(a: Counter, b: Counter, c: Counter) extends Counter { @@ -93,6 +130,7 @@ object BroadcastCounter { b.incr(delta) c.incr(delta) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata, c.metadata)) } private class Four(a: Counter, b: Counter, c: Counter, d: Counter) extends Counter { @@ -102,10 +140,12 @@ object BroadcastCounter { c.incr(delta) d.incr(delta) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata, c.metadata, d.metadata)) } private class N(counters: Seq[Counter]) extends Counter { def incr(delta: Long): Unit = { counters.foreach(_.incr(delta)) } + def metadata: Metadata = MultiMetadata(counters.map(_.metadata)) } } @@ -126,6 +166,7 @@ object BroadcastStat { private object NullStat extends Stat { def add(value: Float): Unit = () + def metadata: Metadata = NoMetadata } private[stats] class Two(a: Stat, b: Stat) extends Stat { @@ -133,6 +174,7 @@ object BroadcastStat { a.add(value) b.add(value) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata)) } private class Three(a: Stat, b: Stat, c: Stat) extends Stat { @@ -141,6 +183,7 @@ object BroadcastStat { b.add(value) c.add(value) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata, c.metadata)) } private class Four(a: Stat, b: Stat, c: Stat, d: Stat) extends Stat { @@ -150,9 +193,11 @@ object BroadcastStat { c.add(value) d.add(value) } + def metadata: Metadata = MultiMetadata(Seq(a.metadata, b.metadata, c.metadata, d.metadata)) } private class N(stats: Seq[Stat]) extends Stat { def add(value: Float): Unit = { stats.foreach(_.add(value)) } + def metadata: Metadata = MultiMetadata(stats.map(_.metadata)) } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/CachedRegex.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/CachedRegex.scala new file mode 100644 index 0000000000..269b22f3e6 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/CachedRegex.scala @@ -0,0 +1,56 @@ +package com.twitter.finagle.stats + +import java.util.concurrent.ConcurrentHashMap +import java.util.function.Consumer +import scala.collection.compat._ +import scala.util.matching.Regex + +/** + * Caches the results of evaluating a given regex + * + * Caches whether the regex matches each of a list of strings. When a map is + * passed to the CachedRegex, it will strip out all of the strings that match + * the regex. + */ +private[stats] class CachedRegex(regex: Regex) + extends (collection.Map[String, Number] => collection.Map[String, Number]) { + // public for testing + // a mapping from a metric key to whether the regex matches it or not + // we allow this to race as long as we're confident that the result is correct + // since a given key will always compute the same value, this is safe to race + val regexMatchCache = new ConcurrentHashMap[String, java.lang.Boolean] + + private[this] val filterFn: String => Boolean = { key => + // we unfortunately need to rely on escape analysis to shave this allocation + // because otherwise java.util.Map#get that returns a null for a plausibly + // primitive type in scala gets casted to a null + // + // the alternative is to first check containsKey and then do a subsequent + // lookup, but that risks an actually potentially bad race condition + // (a potential false negative) + val cached = regexMatchCache.get(key) + val matched = if (cached ne null) { + Boolean.unbox(cached) + } else { + val result = regex.pattern.matcher(key).matches() + // adds new entries + regexMatchCache.put(key, result) + result + } + !matched + } + + def apply(samples: collection.Map[String, Number]): collection.Map[String, Number] = { + // don't really want this to be executed in parallel + regexMatchCache.forEachKey( + Long.MaxValue, + new Consumer[String] { + def accept(key: String): Unit = { + if (!samples.contains(key)) { + regexMatchCache.remove(key) + } + } + }) + samples.view.filterKeys(filterFn).toMap + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/Counter.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/Counter.scala index 1a8286f54f..c3b9034912 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/Counter.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/Counter.scala @@ -7,4 +7,5 @@ package com.twitter.finagle.stats trait Counter { def incr(delta: Long): Unit final def incr(): Unit = { incr(1) } + def metadata: Metadata } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/CumulativeGauge.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/CumulativeGauge.scala index 5473b51b0d..dce44e1605 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/CumulativeGauge.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/CumulativeGauge.scala @@ -1,13 +1,14 @@ package com.twitter.finagle.stats import com.github.benmanes.caffeine.cache.{Cache, Caffeine, RemovalCause, RemovalListener} +import com.twitter.finagle.stats.MetricBuilder.GaugeType import com.twitter.util.lint._ import java.lang.{Boolean => JBoolean} import java.util.concurrent.{ConcurrentHashMap, Executor, ForkJoinPool} import java.util.concurrent.atomic.AtomicInteger -import java.util.concurrent.locks.StampedLock import java.util.function.{Function => JFunction} -import scala.collection.JavaConverters._ +import scala.annotation.tailrec +import scala.jdk.CollectionConverters._ /** * `CumulativeGauge` provides a [[Gauge gauge]] that is composed of the (addition) @@ -22,7 +23,7 @@ private[finagle] abstract class CumulativeGauge(executor: Executor) { self => // uses the default `Executor` for the cache. def this() = this(ForkJoinPool.commonPool()) - private[this] class UnderlyingGauge(val f: () => Float) extends Gauge { + private[this] class UnderlyingGauge(val f: () => Float, val metadata: Metadata) extends Gauge { def remove(): Unit = self.remove(this) } @@ -66,8 +67,8 @@ private[finagle] abstract class CumulativeGauge(executor: Executor) { self => /** * Returns a gauge unless it is dead, in which case it returns null. */ - def addGauge(f: => Float): Gauge = { - val underlyingGauge = new UnderlyingGauge(() => f) + def addGauge(f: => Float, metadata: Metadata): Gauge = { + val underlyingGauge = new UnderlyingGauge(() => f, metadata) refs.put(underlyingGauge, JBoolean.TRUE) if (register()) underlyingGauge else null @@ -110,93 +111,70 @@ trait StatsReceiverWithCumulativeGauges extends StatsReceiver { self => "of Gauges. Indicative of a leak or code registering the same gauge more " + s"often than expected. (For $toString)" ) { - val largeCgs = gauges.asScala.flatMap { - case (ks, cg) => - if (cg.totalSize >= 10000) Some(ks -> cg.totalSize) - else None - } - if (largeCgs.isEmpty) { - Nil - } else { - largeCgs.map { - case (ks, size) => - Issue(ks.mkString("/") + "=" + size) - }.toSeq - } + gauges.asScala.collect { + case (ks, cg) if (cg.totalSize >= 10000) => Issue(ks.mkString("/") + "=" + cg.totalSize) + }.toSeq } } /** - * The StatsReceiver implements these. They provide the cumulated + * The StatsReceiver implements these. They provide the cumulative * gauges. */ - protected[this] def registerGauge(verbosity: Verbosity, name: Seq[String], f: => Float): Unit - protected[this] def deregisterGauge(name: Seq[String]): Unit - - // To avoid creating a new function on each call to `whenGaugeNotPresent`, we cache these - // functions for [[Verbosity.Default]] and [[Verbosity.Debug]], the most common cases. - private[this] val whenDefaultNotPresent = whenNotPresent(Verbosity.Default) - private[this] val whenDebugNotPresent = whenNotPresent(Verbosity.Debug) - - private[this] def getWhenNotPresent(verbosity: Verbosity) = verbosity match { - case Verbosity.Default => whenDefaultNotPresent - case Verbosity.Debug => whenDebugNotPresent - case _ => whenNotPresent(verbosity) - } - - private[this] def whenNotPresent(verbosity: Verbosity) = new JFunction[Seq[String], CumulativeGauge] { - def apply(key: Seq[String]): CumulativeGauge = new CumulativeGauge { - self.registerGauge(verbosity, key, getValue) - - // we grab the read lock on increment and decrement, which can race each - // other. we only grab the write lock when we tear down the cumulative - // gauge, so that we can ensure no more increments happen after death. - private[this] val lock = new StampedLock() - private[this] val registers = new AtomicInteger(0) - - // protected by `lock` - private[this] var dead = false - - def register(): Boolean = { - val stamp = lock.readLock() - try { - if (dead) return false - registers.incrementAndGet() - } finally { - lock.unlockRead(stamp) + protected[this] def registerGauge(metricBuilder: MetricBuilder, f: => Float): Unit + protected[this] def deregisterGauge(metricBuilder: MetricBuilder): Unit + + private[this] def getWhenNotPresent(metricBuilder: MetricBuilder) = whenNotPresent(metricBuilder) + + /** The executor that will be used for expiring gauges */ + def executor: Executor = ForkJoinPool.commonPool() + + private[this] def whenNotPresent(metricBuilder: MetricBuilder) = + new JFunction[Seq[String], CumulativeGauge] { + def apply(key: Seq[String]): CumulativeGauge = new CumulativeGauge(executor) { + self.registerGauge(metricBuilder, getValue) + + // The number of registers starts at `0` because every new gauge will cause a + // registration, even the one that created the `CumulativeGauge`. Because 0 is + // a valid value to increment from, we use `-1` to signal the closed state where + // no more registration activity is possible. + private[this] val registers = new AtomicInteger(0) + + @tailrec def register(): Boolean = { + val c = registers.get + if (c == -1) false + else if (!registers.compareAndSet(c, c + 1)) register() + else true } - true - } - def deregister(): Unit = { - val readStamp = lock.readLock() - if (registers.decrementAndGet() == 0) { - // we could skip a cas by using tryConvertToWriteLock, but it makes - // the code much uglier. - lock.unlockRead(readStamp) - val writeStamp = lock.writeLock() - - try { - if (!dead && registers.get == 0) { - dead = true + @tailrec def deregister(): Unit = registers.get match { + case -1 => () // Already closed + case 1 => + // Attempt to transition to the closed state + if (!registers.compareAndSet(1, -1)) { + deregister() // lost a race so try again + } else { + // Do the cleanup. gauges.remove(key) - self.deregisterGauge(key) + self.deregisterGauge(metricBuilder) + } + + case c => + // Just a normal decrement + if (!registers.compareAndSet(c, c - 1)) { + deregister() // lost a race so try again } - } finally { - lock.unlockWrite(writeStamp) - } - } else { - lock.unlockRead(readStamp) } } } - } - def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge = { + def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + validateMetricType(metricBuilder, GaugeType) var gauge: Gauge = null while (gauge == null) { - val cumulativeGauge = gauges.computeIfAbsent(name, getWhenNotPresent(verbosity)) - gauge = cumulativeGauge.addGauge(f) + val cumulativeGauge = + gauges.computeIfAbsent(metricBuilder.name, getWhenNotPresent(metricBuilder)) + gauge = cumulativeGauge.addGauge(f, metricBuilder) } gauge } @@ -214,5 +192,4 @@ trait StatsReceiverWithCumulativeGauges extends StatsReceiver { self => if (cumulativeGauge == null) 0 else cumulativeGauge.size } - } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/DenylistStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/DenylistStatsReceiver.scala new file mode 100644 index 0000000000..3708a59e0f --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/DenylistStatsReceiver.scala @@ -0,0 +1,61 @@ +package com.twitter.finagle.stats + +object DenylistStatsReceiver { + + /** + * Creates a DenyListStatsReceiver based on a PartialFunction[Seq[String], Boolean]. + * + * @param underlying a base [[StatsReceiver]]. + * @param pf a PartialFunction that returns true for metrics which should be denylisted, + * and false for metrics which should be recorded. + * If pf is undefined for a given metric name, + * then the metric will NOT be recorded. + */ + def orElseDenied( + underlying: StatsReceiver, + pf: PartialFunction[Seq[String], Boolean] + ): StatsReceiver = + new DenylistStatsReceiver(underlying, pf.orElse { case _ => true }) + + /** + * Creates a DenyListStatsReceiver based on a PartialFunction[Seq[String], Boolean]. + * + * @param underlying a base [[StatsReceiver]]. + * @param pf a PartialFunction that returns true for metrics which should be denylisted, + * and false for metrics which should be recorded. + * If pf is undefined for a given metric name, + * then the metric WILL be recorded. + */ + def orElseAdmitted( + underlying: StatsReceiver, + pf: PartialFunction[Seq[String], Boolean] + ): StatsReceiver = + new DenylistStatsReceiver(underlying, pf.orElse { case _ => false }) +} + +/** + * A denylisting [[StatsReceiver]]. If the name for a metric is found to be + * denylisted, nothing is recorded. + * + * @param self a base [[StatsReceiver]], used for metrics that aren't + * denylisted + * @param denylisted a predicate that reads a name and returns true to + * denylist, and false to let it pass through + */ +class DenylistStatsReceiver(protected val self: StatsReceiver, denylisted: Seq[String] => Boolean) + extends StatsReceiverProxy { + + override def counter(metricBuilder: MetricBuilder) = + getStatsReceiver(metricBuilder.name).counter(metricBuilder) + + override def stat(metricBuilder: MetricBuilder) = + getStatsReceiver(metricBuilder.name).stat(metricBuilder) + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float) = + getStatsReceiver(metricBuilder.name).addGauge(metricBuilder)(f) + + private[this] def getStatsReceiver(name: Seq[String]): StatsReceiver = + if (denylisted(name)) NullStatsReceiver else self + + override def toString: String = s"DenylistStatsReceiver($self)" +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/ExceptionStatsHandler.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/ExceptionStatsHandler.scala index 144d975096..bd2d2b7033 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/ExceptionStatsHandler.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/ExceptionStatsHandler.scala @@ -10,7 +10,7 @@ import com.twitter.util.Throwables * @see [[Null]] for a no-op handler. */ object ExceptionStatsHandler { - private[stats] val Failures = "failures" + private[finagle] val Failures = "failures" private[stats] val SourcedFailures = "sourcedfailures" /** @@ -41,11 +41,7 @@ object ExceptionStatsHandler { Seq(exceptionChain, Nil) } - labels.flatMap { prefix => - suffixes.map { suffix => - prefix ++ suffix - } - } + labels.flatMap { prefix => suffixes.map { suffix => prefix ++ suffix } } } } @@ -67,8 +63,8 @@ trait ExceptionStatsHandler { class CategorizingExceptionStatsHandler( categorizer: Throwable => Option[String] = _ => None, sourceFunction: Throwable => Option[String] = _ => None, - rollup: Boolean = true -) extends ExceptionStatsHandler { + rollup: Boolean = true) + extends ExceptionStatsHandler { import ExceptionStatsHandler._ private[this] val underlying: ExceptionStatsHandler = { @@ -115,12 +111,18 @@ class CategorizingExceptionStatsHandler( * failures/restartable/com.twitter.finagle.Failure/com.twitter.finagle.CancelledConnectionException: 1 * failures/restartable/com.twitter.finagle.Failure/com.twitter.finagle.CancelledConnectionException/java.net.ConnectException: 1 * + * `f1` is also recorded dimensionally as: + * failures{interrupted=true, restartable=true, exception=java.net.ConnectException} + * * `f2` is recorded as: * failures * failures/com.twitter.finagle.IndividualRequestTimeoutException * sourcedfailures/myservice * sourcedfailures/myservice/com.twitter.finagle.IndividualRequestTimeoutException * + * `f2` is also recorded dimensionally as: + * failures{interrupted=true, source=myservice, restartable=true, exception=com.twitter.finagle.IndividualRequestTimeoutException} + * * @param mkLabel label prefix, default to 'failures' * @param mkFlags extracting flags if the Throwable is a Failure * @param mkSource extracting source when configured @@ -131,8 +133,8 @@ private[finagle] class MultiCategorizingExceptionStatsHandler( mkLabel: Throwable => String = _ => ExceptionStatsHandler.Failures, mkFlags: Throwable => Set[String] = _ => Set.empty, mkSource: Throwable => Option[String] = _ => None, - rollup: Boolean = true -) extends ExceptionStatsHandler { + rollup: Boolean = true) + extends ExceptionStatsHandler { import ExceptionStatsHandler._ def record(statsReceiver: StatsReceiver, t: Throwable): Unit = { @@ -143,16 +145,32 @@ private[finagle] class MultiCategorizingExceptionStatsHandler( if (flags.isEmpty) Seq(Seq(parentLabel)) else flags.toSeq.map(Seq(parentLabel, _)) - val labels: Seq[Seq[String]] = mkSource(t) match { + val maybeSource = mkSource(t) + val labels: Seq[Seq[String]] = maybeSource match { case Some(service) => flagLabels :+ Seq(SourcedFailures, service) case None => flagLabels } val paths: Seq[Seq[String]] = statPaths(t, labels, rollup) - if (flags.nonEmpty) statsReceiver.counter(parentLabel).incr() - paths.foreach { path => - statsReceiver.counter(path: _*).incr() + val dimensionalLabels = Map("exception" -> Throwables.RootCause.nested(t).getClass.getName) ++ + flags.map { flag => flag -> "true" } ++ + maybeSource.map("source" -> _).toSeq + + val dimensionalMetric = MetricBuilder.forCounter + .withName(parentLabel) + .withLabels(dimensionalLabels) + .withDimensionalSupport + + statsReceiver.counter(dimensionalMetric).incr() + paths.foreach { + case Seq(comp) if comp == parentLabel => + // The "flat" counter is used to build up the rollup hierarchies but since we also use it + // to hold the dimensional metric, we skip it when recording hierarchical metrics. + case path => + statsReceiver + .counter(MetricBuilder.forCounter.withHierarchicalOnly.withName(path: _*)) + .incr() } } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/Gauge.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/Gauge.scala index a6f8ea43ad..ed3126254d 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/Gauge.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/Gauge.scala @@ -6,4 +6,5 @@ package com.twitter.finagle.stats */ trait Gauge { def remove(): Unit + def metadata: Metadata } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramCounter.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramCounter.scala new file mode 100644 index 0000000000..719f3484ee --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramCounter.scala @@ -0,0 +1,106 @@ +package com.twitter.finagle.stats + +import com.twitter.conversions.DurationOps._ +import com.twitter.util.Closable +import com.twitter.util.Duration +import com.twitter.util.Future +import com.twitter.util.Time +import com.twitter.util.Timer +import java.util.concurrent.ConcurrentHashMap +import java.util.concurrent.atomic.LongAdder +import scala.collection.JavaConverters._ + +private[twitter] sealed abstract class StatsFrequency(val frequency: Duration) { + def suffix: String +} + +private[twitter] object StatsFrequency { + case object HundredMilliSecondly extends StatsFrequency(100.millis) { + override def suffix = "hundredMilliSecondly" + } +} + +/** + * Class for creating [[HistogramCounter]]s. It is expected that there be one [[HistogramCounterFactory]] + * per process -- otherwise we will schedule multiple timer tasks for aggregating the counter into + * a stat, and there can be multiple aggregations for a single stat which may produce unexpected + * results. + */ +private[twitter] class HistogramCounterFactory(timer: Timer, nowMs: () => Long) extends Closable { + + @volatile private[this] var closed = false + + private[this] val frequencyToStats: Map[ + StatsFrequency, + ConcurrentHashMap[Stat, HistogramCounter] + ] = Map( + StatsFrequency.HundredMilliSecondly -> new ConcurrentHashMap[Stat, HistogramCounter] + ) + + frequencyToStats.map { + case (statsFrequency, statToCounter) => + timer.doLater(statsFrequency.frequency)(recordStatsForCounters(statsFrequency, statToCounter)) + } + + def apply( + name: Seq[String], + frequency: StatsFrequency, + statsReceiver: StatsReceiver + ): HistogramCounter = { + val stat = statsReceiver.stat(normalizeName(name) :+ frequency.suffix: _*) + val histogramCounter = new HistogramCounter(stat, nowMs, frequency.frequency.inMillis) + val existing = frequencyToStats(frequency).putIfAbsent(stat, histogramCounter) + if (existing == null) { + histogramCounter + } else { + existing + } + } + + override def close(deadline: Time): Future[Unit] = { + closed = true + Future.Done + } + + private[this] def recordStatsForCounters( + statsFrequency: StatsFrequency, + statToCounter: ConcurrentHashMap[Stat, HistogramCounter] + ): Unit = { + statToCounter.values().asScala.foreach { counter => + counter.recordAndReset() + } + if (!closed) { + timer.doLater(statsFrequency.frequency)(recordStatsForCounters(statsFrequency, statToCounter)) + } + } + + private[this] def normalizeName(name: Seq[String]): Seq[String] = { + if (name.forall(!_.contains("/"))) { + name + } else { + name.map(_.split("/")).flatten + } + } +} + +private[stats] class HistogramCounter(stat: Stat, nowMs: () => Long, windowSizeMs: Long) { + private[this] val counter: LongAdder = new LongAdder + @volatile private[this] var lastRecordAndResetMs = nowMs() + + private[stats] def recordAndReset(): Unit = { + val count = counter.sumThenReset() + val now = nowMs() + val elapsed = Math.max(0, now - lastRecordAndResetMs) + val elapsedWindows = elapsed.toFloat / windowSizeMs + stat.add(count / elapsedWindows) + lastRecordAndResetMs = now + } + + def incr(delta: Long): Unit = { + counter.add(delta) + } + + def incr(): Unit = { + counter.increment() + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramFormatter.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramFormatter.scala new file mode 100644 index 0000000000..c5cb743e2e --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/HistogramFormatter.scala @@ -0,0 +1,104 @@ +package com.twitter.finagle.stats + +import com.twitter.app.GlobalFlag + +object format + extends GlobalFlag[String]( + "commonsmetrics", + "Format style for metric names (ostrich|commonsmetrics|commonsstats)" + ) { + private[stats] val Ostrich = "ostrich" + private[stats] val CommonsMetrics = "commonsmetrics" + private[stats] val CommonsStats = "commonsstats" +} + +private[twitter] object HistogramFormatter { + def default: HistogramFormatter = + format() match { + case format.Ostrich => Ostrich + case format.CommonsMetrics => CommonsMetrics + case format.CommonsStats => CommonsStats + } + + /** + * The default behavior for formatting as done by Commons Metrics. + * + * See Commons Metrics' `Metrics.sample()`. + */ + object CommonsMetrics extends HistogramFormatter { + def labelPercentile(p: Double): String = { + // this has a strange quirk that p999 gets formatted as p9990 + // Round for precision issues; e.g. 0.9998999... converts to "p9998" with a direct int cast. + val gname: String = "p" + (p * 10000).round + if (3 < gname.length && ("00" == gname.substring(3))) { + gname.substring(0, 3) + } else { + gname + } + } + + override def labelMin: String = "min" + + override def labelMax: String = "max" + + override def labelAverage: String = "avg" + } + + /** + * Replicates the behavior for formatting Ostrich stats. + * + * See Ostrich's `Distribution.toMap`. + */ + object Ostrich extends HistogramFormatter { + def labelPercentile(p: Double): String = { + p match { + case 0.5d => "p50" + case 0.9d => "p90" + case 0.95d => "p95" + case 0.99d => "p99" + case 0.999d => "p999" + case 0.9999d => "p9999" + case _ => + val padded = (p * 10000).toInt + s"p$padded" + } + } + + override def labelMin: String = "minimum" + + override def labelMax: String = "maximum" + + override def labelAverage: String = "average" + } + + /** + * Replicates the behavior for formatting Commons Stats stats. + * + * See Commons Stats' `Stats.getVariables()`. + */ + object CommonsStats extends HistogramFormatter { + def labelPercentile(p: Double): String = + s"${p * 100}_percentile".replace(".", "_") + + override def labelMin: String = "min" + + override def labelMax: String = "max" + + override def labelAverage: String = "avg" + } +} + +sealed trait HistogramFormatter { + def labelPercentile(p: Double): String + + def labelMin: String = "min" + + def labelMax: String = "max" + + def labelAverage: String = "avg" + + def labelCount: String = "count" + + def labelSum: String = "sum" + +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala index dd106b5bd3..587a4b90c1 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/InMemoryStatsReceiver.scala @@ -1,9 +1,22 @@ package com.twitter.finagle.stats +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.exp.ExpressionSchema.ExpressionCollisionException +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.finagle.stats.exp.ExpressionSchemaKey +import com.twitter.util.Throw +import com.twitter.util.Try import java.io.PrintStream +import java.util.Locale import java.util.concurrent.ConcurrentHashMap -import scala.collection.JavaConverters._ -import scala.collection.{SortedMap, mutable} +import scala.annotation.varargs +import scala.collection.compat._ +import scala.collection.SortedMap +import scala.collection.mutable +import scala.collection.{Map => scalaMap} +import scala.jdk.CollectionConverters._ object InMemoryStatsReceiver { private[stats] implicit class RichMap[K, V](val self: mutable.Map[K, V]) { @@ -65,100 +78,141 @@ class InMemoryStatsReceiver extends StatsReceiver with WithHistogramDetails { val gauges: mutable.Map[Seq[String], () => Float] = new ConcurrentHashMap[Seq[String], () => Float]().asScala + val schemas: mutable.Map[Seq[String], MetricBuilder] = + new ConcurrentHashMap[Seq[String], MetricBuilder]().asScala + + val expressions: mutable.Map[ExpressionSchemaKey, ExpressionSchema] = + new ConcurrentHashMap[ExpressionSchemaKey, ExpressionSchema]().asScala + + // duplicate metric name -> duplication times + val duplicatedMetrics: mutable.Map[Seq[String], Long] = + new ConcurrentHashMap[Seq[String], Long]().asScala + + private def recordDup(metricBuilder: MetricBuilder): Unit = { + if (!duplicatedMetrics.contains(metricBuilder.name)) { + duplicatedMetrics(metricBuilder.name) = 1 + } else { + duplicatedMetrics(metricBuilder.name) = duplicatedMetrics(metricBuilder.name) + 1 + } + } + + @varargs override def counter(name: String*): ReadableCounter = - counter(Verbosity.Default, name: _*) + counter(MetricBuilder.forCounter.withName(name: _*)) /** * Creates a [[ReadableCounter]] of the given `name`. */ - def counter(v: Verbosity, name: String*): ReadableCounter = + def counter(metricBuilder: MetricBuilder): ReadableCounter = { + validateMetricType(metricBuilder, CounterType) new ReadableCounter { - verbosity += name -> v + verbosity += metricBuilder.name -> metricBuilder.verbosity // eagerly initialize counters.synchronized { - if (!counters.contains(name)) { - counters(name) = 0 + if (!counters.contains(metricBuilder.name)) { + counters(metricBuilder.name) = 0 + schemas(metricBuilder.name) = metricBuilder + } else { + recordDup(metricBuilder) } } def incr(delta: Long): Unit = counters.synchronized { val oldValue = apply() - counters(name) = oldValue + delta + counters(metricBuilder.name) = oldValue + delta } - def apply(): Long = counters.getOrElse(name, 0) + def apply(): Long = counters.getOrElse(metricBuilder.name, 0) + + def metadata: Metadata = metricBuilder override def toString: String = - s"Counter(${name.mkString("/")}=${apply()})" + s"Counter(${metricBuilder.name.mkString("/")}=${apply()})" } + } - override def stat(name: String*): ReadableStat = stat(Verbosity.Default, name: _*) + @varargs + override def stat(name: String*): ReadableStat = + stat(MetricBuilder.forStat.withName(name: _*)) /** * Creates a [[ReadableStat]] of the given `name`. */ - def stat(v: Verbosity, name: String*): ReadableStat = + def stat(metricBuilder: MetricBuilder): ReadableStat = { + validateMetricType(metricBuilder, HistogramType) new ReadableStat { - verbosity += name -> v + verbosity += metricBuilder.name -> metricBuilder.verbosity // eagerly initialize stats.synchronized { - if (!stats.contains(name)) { - stats(name) = Nil + if (!stats.contains(metricBuilder.name)) { + stats(metricBuilder.name) = Nil + schemas(metricBuilder.name) = metricBuilder + } else { + recordDup(metricBuilder) } } def add(value: Float): Unit = stats.synchronized { val oldValue = apply() - stats(name) = oldValue :+ value + stats(metricBuilder.name) = oldValue :+ value } - def apply(): Seq[Float] = stats.getOrElse(name, Seq.empty) + def apply(): Seq[Float] = stats.getOrElse(metricBuilder.name, Seq.empty) + + def metadata: Metadata = metricBuilder override def toString: String = { val vals = apply() - s"Stat(${name.mkString("/")}=${statValuesToStr(vals)})" + s"Stat(${metricBuilder.name.mkString("/")}=${statValuesToStr(vals)})" } } + } /** * Creates a [[Gauge]] of the given `name`. */ - def addGauge(v: Verbosity, name: String*)(f: => Float): Gauge = + def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + validateMetricType(metricBuilder, GaugeType) new Gauge { - gauges += name -> (() => f) - verbosity += name -> v + gauges += metricBuilder.name -> (() => f) + schemas += metricBuilder.name -> metricBuilder + verbosity += metricBuilder.name -> metricBuilder.verbosity def remove(): Unit = { - gauges -= name + gauges -= metricBuilder.name + schemas -= metricBuilder.name } + def metadata: Metadata = metricBuilder + override def toString: String = { // avoid holding a reference to `f` - val current = gauges.get(name) match { + val current = gauges.get(metricBuilder.name) match { case Some(fn) => fn() case None => -0.0f } - s"Gauge(${name.mkString("/")}=$current)" + s"Gauge(${metricBuilder.name.mkString("/")}=$current)" } } + } override def toString: String = "InMemoryStatsReceiver" /** - * Dumps this in-memory stats receiver to the given [[PrintStream]]. - * @param p the [[PrintStream]] to which to write in-memory values. + * Dumps this in-memory stats receiver to the given `PrintStream`. + * @param p the `PrintStream` to which to write in-memory values. */ def print(p: PrintStream): Unit = { print(p, includeHeaders = false) } /** - * Dumps this in-memory stats receiver to the given [[PrintStream]]. - * @param p the [[PrintStream]] to which to write in-memory values. + * Dumps this in-memory stats receiver to the given `PrintStream`. + * @param p the `PrintStream` to which to write in-memory values. * @param includeHeaders optionally include printing underlines headers for the different types * of stats printed, e.g., "Counters:", "Gauges:", "Stats;" */ @@ -172,30 +226,44 @@ class InMemoryStatsReceiver extends StatsReceiver with WithHistogramDetails { p.println("---------") } for ((k, v) <- sortedCounters) - p.print(f"$k%s $v%d\n") + p.println(f"$k%s $v%d") if (includeHeaders && sortedGauges.nonEmpty) { p.println("\nGauges:") p.println("-------") } - for ((k, g) <- sortedGauges) - p.print(f"$k%s ${ g() }%f\n") + for ((k, g) <- sortedGauges) { + p.println("%s %f".formatLocal(Locale.US, k, g())) + } if (includeHeaders && sortedStats.nonEmpty) { p.println("\nStats:") p.println("------") } for ((k, s) <- sortedStats if s.nonEmpty) { - p.print(f"$k%s ${ s.sum / s.size }%f ${statValuesToStr(s)}\n") + p.println("%s %f %s".formatLocal(Locale.US, k, s.sum / s.size, statValuesToStr(s))) } } /** - * Clears all registered counters, gauges and stats. + * Dumps this in-memory stats receiver's MetricMetadata to the given `PrintStream`. + * @param p the `PrintStream` to which to write in-memory metadata. + */ + def printSchemas(p: PrintStream): Unit = { + val sortedSchemas = schemas.mapKeys(_.mkString("/")).toSortedMap + for ((k, schema) <- sortedSchemas) { + p.println(s"$k $schema") + } + } + + /** + * Clears all registered counters, gauges, stats, expressions, and their metadata. * @note this is not atomic. If new metrics are added while this method is executing, those metrics may remain. */ def clear(): Unit = { counters.clear() stats.clear() gauges.clear() + schemas.clear() + expressions.clear() } private[this] def toHistogramDetail(addedValues: Seq[Float]): HistogramDetail = { @@ -208,14 +276,10 @@ class InMemoryStatsReceiver extends StatsReceiver with WithHistogramDetails { new HistogramDetail { def counts: Seq[BucketAndCount] = { addedValues - .map { x => - nearestPosInt(x) - } - .groupBy(identity) - .mapValues(_.size) + .groupBy(nearestPosInt) + .map { case (k, vs) => BucketAndCount(k, k + 1, vs.size) } .toSeq - .sortWith(_._1 < _._1) - .map { case (k, v) => BucketAndCount(k, k + 1, v) } + .sortBy(_.lowerLimit) } } } @@ -223,6 +287,29 @@ class InMemoryStatsReceiver extends StatsReceiver with WithHistogramDetails { def histogramDetails: Map[String, HistogramDetail] = stats.toMap.map { case (k, v) => (k.mkString("/"), toHistogramDetail(v)) } + + /** + * Designed to match the behavior of Metrics::registerExpression(). + */ + override def registerExpression(schema: ExpressionSchema): Try[Unit] = { + if (expressions.contains(schema.schemaKey())) { + Throw( + ExpressionCollisionException( + s"An expression with the key ${schema.schemaKey()} had already been defined.")) + } else Try { expressions.put(schema.schemaKey(), schema) } + } + + /** + * Retrieves all expressions with a given label key and value + * + * @return a Map of expressions ([[ExpressionSchemaKey]] -> [[ExpressionSchema]]) + */ + def getAllExpressionsWithLabel( + key: String, + value: String + ): scalaMap[ExpressionSchemaKey, ExpressionSchema] = + expressions.view.filterKeys(k => k.labels.get(key) == Some(value)).toMap + } /** diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/KeyIndicatorMarkingStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/KeyIndicatorMarkingStatsReceiver.scala new file mode 100644 index 0000000000..e9d7a8989a --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/KeyIndicatorMarkingStatsReceiver.scala @@ -0,0 +1,23 @@ +package com.twitter.finagle.stats + +/** + * A StatsReceiver proxy that configures all counters, stats, and gauges marked as KeyIndicators. + * This is used with [[ReadableIndicator]]. + * + * @param self The underlying StatsReceiver to which modified schemas are passed + */ +private[finagle] class KeyIndicatorMarkingStatsReceiver(protected val self: StatsReceiver) + extends StatsReceiverProxy { + + override def counter(metricBuilder: MetricBuilder): Counter = { + self.counter(metricBuilder.withKeyIndicator()) + } + + override def stat(metricBuilder: MetricBuilder): Stat = { + self.stat(metricBuilder.withKeyIndicator()) + } + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + self.addGauge(metricBuilder.withKeyIndicator())(f) + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/LazyStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/LazyStatsReceiver.scala new file mode 100644 index 0000000000..f87c3f4bc8 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/LazyStatsReceiver.scala @@ -0,0 +1,41 @@ +package com.twitter.finagle.stats + +import com.twitter.finagle.stats.MetricBuilder.{CounterType, HistogramType} + +/** + * Wraps an underlying [[StatsReceiver]] to ensure that derived counters and + * stats will not start exporting metrics until `incr` or `add` is first called + * on them. + * + * This should be used when integrating with tools that create metrics eagerly, + * but you don't know whether you're going to actually use those metrics or not. + * One example might be if you're speaking to a remote peer that exposes many + * endpoints, and you eagerly create metrics for all of those endpoints, but + * aren't going to use all of the different methods. + * + * We don't apply this very widely automatically--it can mess with caching, and + * adds an extra allocation when you construct a new counter or stat, so please + * be judicious when using it. + * + * @note does not change the way gauges are used, since there isn't a way of + * modeling whether a gauge is "used" or not. + */ +final class LazyStatsReceiver(val self: StatsReceiver) extends StatsReceiverProxy { + override def counter(metricBuilder: MetricBuilder): Counter = { + validateMetricType(metricBuilder, CounterType) + new Counter { + private[this] lazy val underlying = self.counter(metricBuilder) + def incr(delta: Long): Unit = underlying.incr(delta) + def metadata: Metadata = metricBuilder + } + } + + override def stat(metricBuilder: MetricBuilder): Stat = { + validateMetricType(metricBuilder, HistogramType) + new Stat { + private[this] lazy val underlying = self.stat(metricBuilder) + def add(value: Float): Unit = underlying.add(value) + def metadata: Metadata = metricBuilder + } + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/LoadedStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/LoadedStatsReceiver.scala new file mode 100644 index 0000000000..a3098a5cd7 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/LoadedStatsReceiver.scala @@ -0,0 +1,31 @@ +package com.twitter.finagle.stats + +import com.twitter.app.LoadService + +/** + * A [[com.twitter.finagle.stats.StatsReceiver]] that loads + * all service-loadable receivers and broadcasts stats to them. + */ +object LoadedStatsReceiver extends StatsReceiverProxy { + + /** + * Mutating this value at runtime after it has been initialized should be done + * with great care. If metrics have been created using the prior + * [[StatsReceiver]], updates to those metrics may not be reflected in the + * [[StatsReceiver]] that replaces it. In addition, histograms created with + * the prior [[StatsReceiver]] will not be available. + */ + @volatile var self: StatsReceiver = BroadcastStatsReceiver(LoadService[StatsReceiver]()) +} + +/** + * A "default" StatsReceiver loaded by the + * [[com.twitter.finagle.util.LoadService]] mechanism. + */ +object DefaultStatsReceiver extends StatsReceiverProxy { + val self: StatsReceiver = LoadedStatsReceiver.label("implementation", "app") + + override def repr: DefaultStatsReceiver.type = this + + def get: StatsReceiver = this +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/MetricBuilder.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricBuilder.scala new file mode 100644 index 0000000000..47a7ddf8b3 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricBuilder.scala @@ -0,0 +1,506 @@ +package com.twitter.finagle.stats + +import com.twitter.app.GlobalFlag +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.CounterishGaugeType +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.MetricBuilder.Identity +import com.twitter.finagle.stats.MetricBuilder.IdentityType +import com.twitter.finagle.stats.MetricBuilder.MetricType +import com.twitter.finagle.stats.MetricBuilder.UnlatchedCounter +import scala.annotation.varargs + +object biasToFull + extends GlobalFlag[Boolean]( + false, + "Bias metrics with an indeterminate identity to be exported through all available exporters" + ) + +/** + * Represents the "role" this service plays with respect to this metric. + * + * Usually either Server (the service is processing a request) or Client (the server has sent a + * request to another service). In many cases, there is no relevant role for a metric, in which + * case NoRole should be used. + */ +sealed abstract class SourceRole private () { + + /** + * Java-friendly helper for accessing the object itself. + */ + def getInstance(): SourceRole = this +} + +object SourceRole { + case object NoRoleSpecified extends SourceRole + case object Client extends SourceRole + case object Server extends SourceRole + + /** + * Java-friendly accessors + */ + def noRoleSpecified(): SourceRole = NoRoleSpecified + + def client(): SourceRole = Client + + def server(): SourceRole = Server +} + +/** + * A format for when histogram is exported. A few formats are supported: + * + * - `NoSummary`: only percentiles (p50, ..., p9999) are exported + * - `ShortSummary`: percentiles and `count`, and `sum` are exported + * - `FullSummary`: percentiles, `count`, `sum`, `avg`, `min`, and `max` are exported + * - `Default`: let the exporting sub-system (usually, TwitterServer) decide + */ +sealed abstract class HistogramFormat private (override val toString: String) + +object HistogramFormat { + case object Default extends HistogramFormat("default") + case object NoSummary extends HistogramFormat("no_summary") + case object ShortSummary extends HistogramFormat("short_summary") + case object FullSummary extends HistogramFormat("full_summary") +} + +/** + * finagle-stats has configurable scope separators. As this package is wrapped by finagle-stats, we + * cannot retrieve it from finagle-stats. Consequently, we created this object so that the + * scope-separator can be passed in for stringification of the MetricBuilder objects. + */ +object metadataScopeSeparator { + @volatile private var separator: String = "/" + + def apply(): String = separator + + private[finagle] def setSeparator(separator: String): Unit = { + this.separator = separator + } +} + +object MetricBuilder { + + val DimensionalNameScopeSeparator: String = "_" + + val NoDescription: String = "No description provided" + + private[this] val EmptyPercentiles = IndexedSeq.empty + + /** + * Construct a MethodBuilder. + * + * @param keyIndicator indicates whether this metric is crucial to this service (ie, an SLO metric) + * @param description human-readable description of a metric's significance + * @param units the unit associated with the metrics value (milliseconds, megabytes, requests, etc) + * @param role whether the service is playing the part of client or server regarding this metric + * @param verbosity see StatsReceiver for details + * @param sourceClass the name of the class which generated this metric (ie, com.twitter.finagle.StatsFilter) + * @param name the full metric name + * @param relativeName the relative metric name which will be appended to the scope of the StatsReceiver prior to long term storage + * @param processPath a universal coordinate for the resource + * @param percentiles used to indicate buckets for histograms, to be set by the StatsReceiver + * @param metricType the type of the metric being constructed + * @param isStandard whether the metric should the standard path separator or defer to the metrics store for the separator + */ + def apply( + keyIndicator: Boolean = false, + description: String = NoDescription, + units: MetricUnit = Unspecified, + role: SourceRole = SourceRole.NoRoleSpecified, + verbosity: Verbosity = Verbosity.Default, + sourceClass: Option[String] = None, + name: Seq[String] = Seq.empty, + relativeName: Seq[String] = Seq.empty, + processPath: Option[String] = None, + histogramFormat: HistogramFormat = HistogramFormat.Default, + percentiles: IndexedSeq[Double] = + EmptyPercentiles, // IndexedSeq.empty creates a new collection on every call + isStandard: Boolean = false, + metricType: MetricType, + ): MetricBuilder = { + new MetricBuilder( + keyIndicator = keyIndicator, + description = description, + units = units, + role = role, + verbosity = verbosity, + sourceClass = sourceClass, + // For now we're not allowing construction w"ith labels. + // They need to be added later via `.withLabels`. + identity = Identity(name, name), + relativeName = relativeName, + processPath = processPath, + histogramFormat = histogramFormat, + percentiles = percentiles, + isStandard = isStandard, + metricType = metricType, + metricUsageHints = Set.empty + ) + } + + /** + * Construct a new `MetricBuilder` + * + * Simpler method for constructing a `MetricBuilder`. This API is a little more friendly for java + * users as they can bootstrap a `MetricBuilder` and then proceed to use with `.with` API's rather + * than being forced to use the mega apply method. + * + * @param metricType the type of the metric being constructed. + */ + def newBuilder(metricType: MetricType): MetricBuilder = + apply(metricType = metricType) + + /** + * Construct a new `MetricBuilder` that will build counters. The result is the same as calling + * `newBuilder(CounterType)`. + */ + val forCounter: MetricBuilder = newBuilder(MetricBuilder.CounterType) + + /** + * Construct a new `MetricBuilder` that will build stats (histograms). The result is the same as + * calling `newBuilder(HistogramType)`. + */ + val forStat: MetricBuilder = newBuilder(MetricBuilder.HistogramType) + + /** + * Construct a new `MetricBuilder` that will build gauges. The result is the same as calling + * `newBuilder(GaugeType)`. + */ + val forGauge: MetricBuilder = newBuilder(MetricBuilder.GaugeType) + + /** + * Indicate the Metric type, [[CounterType]] will create a [[Counter]], + * [[GaugeType]] will create a standard [[Gauge]], and [[HistogramType]] will create a [[Stat]]. + * + * @note [[CounterishGaugeType]] will also create a [[Gauge]], but specifically one that + * models a [[Counter]]. + */ + sealed trait MetricType { + def toPrometheusString: String + def toJsonString: String + } + case object CounterType extends MetricType { + def toJsonString: String = "counter" + def toPrometheusString: String = toJsonString + } + case object CounterishGaugeType extends MetricType { + def toJsonString: String = "counterish_gauge" + def toPrometheusString: String = "gauge" + } + case object GaugeType extends MetricType { + def toJsonString: String = "gauge" + def toPrometheusString: String = toJsonString + } + case object HistogramType extends MetricType { + def toJsonString: String = "histogram" + def toPrometheusString: String = "summary" + } + + /** Counter type that will be always unlatched */ + case object UnlatchedCounter extends MetricType { + def toJsonString: String = "unlatched_counter" + def toPrometheusString: String = "counter" + } + + /** Representation of how an [[Identity]]s name is defined */ + sealed abstract class IdentityType private () { + + /** + * Bias this [[IdentityType]] to the specified type. + * + * @param to in the case that this is a `NonDeterminate` identity type, what the result should be. + */ + final def bias(to: IdentityType.ResolvedIdentityType): IdentityType.ResolvedIdentityType = + this match { + case IdentityType.NonDeterminate => to + case IdentityType.Full => IdentityType.Full + case IdentityType.HierarchicalOnly => IdentityType.HierarchicalOnly + } + } + + object IdentityType { + + /** + * Resolves an [[IdentityType]] to the default [[ResolvedIdentityType]] + * + * A centralized place to resolve the ambiguity of the [[NonDeterminate]] identity type. + * We often need to know if a metric has dimensional support or not such as when emitting + * a metric. This function offers a central place where that decision is made so that in + * the future it can be changed in a central location or even put behind a flag. + * + * The current behavior is to default to [[IdentityType.HierarchicalOnly]] + */ + private[twitter] def toResolvedIdentityType(identityType: IdentityType): ResolvedIdentityType = + identityType.bias(if (biasToFull()) IdentityType.Full else IdentityType.HierarchicalOnly) + + /** An [[IdentityType]] that cannot be [[NonDeterminate]] */ + sealed abstract class ResolvedIdentityType extends IdentityType + + /** + * An [[IdentityType]] which has a name that is currently valid both dimensionally and + * hierarchically, but leaves it up to the [[StatsReceiver]] to decide whether to represent + * it with one or both names. + */ + case object NonDeterminate extends IdentityType + + /** + * An [[IdentityType]] that has means the associated metric [[Identity]] has been + * explicitly marked to not be valid dimensionally + */ + case object HierarchicalOnly extends ResolvedIdentityType + + /** + * An [[IdentityType]] which means that the associated metric [[Identity]] is valid + * both dimensionally and hierarchically + */ + case object Full extends ResolvedIdentityType + } + + /** + * Experimental way to represent both hierarchical and dimensional identity information. + * + * @note this type is expected to undergo some churn as dimensional metric support matures. + */ + final case class Identity( + hierarchicalName: Seq[String], + dimensionalName: Seq[String], + labels: Map[String, String] = Map.empty, + identityType: IdentityType = IdentityType.NonDeterminate) { + override def toString: String = { + val hname = hierarchicalName.mkString(metadataScopeSeparator()) + val dname = dimensionalName.mkString(DimensionalNameScopeSeparator) + s"Identity($hname, $dname, $labels, $identityType)" + } + } +} + +sealed trait Metadata { + + /** + * Extract the MetricBuilder from Metadata + * + * Will return `None` if it's `NoMetadata` + */ + def toMetricBuilder: Option[MetricBuilder] = this match { + case metricBuilder: MetricBuilder => Some(metricBuilder) + case NoMetadata => None + case MultiMetadata(schemas) => + schemas.find(_ != NoMetadata).flatMap(_.toMetricBuilder) + } +} + +case object NoMetadata extends Metadata + +case class MultiMetadata(schemas: Seq[Metadata]) extends Metadata + +/** + * A builder class used to configure settings and metadata for metrics prior to instantiating them. + * Calling any of the three build methods (counter, gauge, or histogram) will cause the metric to be + * instantiated in the underlying StatsReceiver. + */ +class MetricBuilder private ( + val keyIndicator: Boolean, + val description: String, + val units: MetricUnit, + val role: SourceRole, + val verbosity: Verbosity, + val sourceClass: Option[String], + val identity: Identity, + val relativeName: Seq[String], + val processPath: Option[String], + // Only persisted and relevant when building histograms. + val histogramFormat: HistogramFormat, + val percentiles: IndexedSeq[Double], + val isStandard: Boolean, + val metricType: MetricType, + val metricUsageHints: Set[MetricUsageHint]) + extends Metadata { + + /** + * This copy method omits statsReceiver and percentiles as arguments because they should be + * set only during the initial creation of a MetricBuilder by a StatsReceiver, which should use + * itself as the value for statsReceiver. + */ + private[this] def copy( + keyIndicator: Boolean = this.keyIndicator, + description: String = this.description, + units: MetricUnit = this.units, + role: SourceRole = this.role, + verbosity: Verbosity = this.verbosity, + sourceClass: Option[String] = this.sourceClass, + identity: Identity = this.identity, + relativeName: Seq[String] = this.relativeName, + processPath: Option[String] = this.processPath, + isStandard: Boolean = this.isStandard, + histogramFormat: HistogramFormat = this.histogramFormat, + percentiles: IndexedSeq[Double] = this.percentiles, + metricType: MetricType = this.metricType, + metricUsageHints: Set[MetricUsageHint] = this.metricUsageHints + ): MetricBuilder = { + new MetricBuilder( + keyIndicator = keyIndicator, + description = description, + units = units, + role = role, + verbosity = verbosity, + sourceClass = sourceClass, + identity = identity, + relativeName = relativeName, + processPath = processPath, + histogramFormat = histogramFormat, + percentiles = percentiles, + isStandard = isStandard, + metricType = metricType, + metricUsageHints = metricUsageHints + ) + } + + private[finagle] def withStandard: MetricBuilder = { + this.copy(isStandard = true) + } + + private[finagle] def withUnlatchedCounter: MetricBuilder = { + require(this.metricType == CounterType, "unable to set a non counter to unlatched counter") + this.copy(metricType = UnlatchedCounter) + } + + def withKeyIndicator(isKeyIndicator: Boolean = true): MetricBuilder = + this.copy(keyIndicator = isKeyIndicator) + + def withDescription(desc: String): MetricBuilder = this.copy(description = desc) + + def withVerbosity(verbosity: Verbosity): MetricBuilder = this.copy(verbosity = verbosity) + + def withSourceClass(sourceClass: Option[String]): MetricBuilder = + this.copy(sourceClass = sourceClass) + + def withIdentifier(processPath: Option[String]): MetricBuilder = + this.copy(processPath = processPath) + + def withUnits(units: MetricUnit): MetricBuilder = this.copy(units = units) + + def withRole(role: SourceRole): MetricBuilder = this.copy(role = role) + + @varargs + def withName(name: String*): MetricBuilder = { + copy(identity = identity.copy(hierarchicalName = name, dimensionalName = name)) + } + + private[twitter] def withIdentity(identity: Identity): MetricBuilder = copy(identity = identity) + + /** + * The hierarchical name of the metric + */ + def name: Seq[String] = identity.hierarchicalName + + // Private for now + private[twitter] def withLabels(labels: Map[String, String]): MetricBuilder = { + copy(identity = identity.copy(labels = labels)) + } + + private[twitter] def withHierarchicalOnly: MetricBuilder = { + if (identity.identityType == IdentityType.HierarchicalOnly) this + else copy(identity = identity.copy(identityType = IdentityType.HierarchicalOnly)) + } + + private[twitter] def withDimensionalSupport: MetricBuilder = { + if (identity.identityType == IdentityType.Full) this + else copy(identity = identity.copy(identityType = IdentityType.Full)) + } + + @varargs + def withRelativeName(relativeName: String*): MetricBuilder = + if (this.relativeName == Seq.empty) this.copy(relativeName = relativeName) else this + + def withHistogramFormat(histogramFormat: HistogramFormat): MetricBuilder = + this.copy(histogramFormat = histogramFormat) + + @varargs + def withPercentiles(percentiles: Double*): MetricBuilder = + this.copy(percentiles = percentiles.toIndexedSeq) + + def withCounterishGauge: MetricBuilder = { + require( + this.metricType == GaugeType, + "Cannot create a CounterishGauge from a Counter or Histogram") + this.copy(metricType = CounterishGaugeType) + } + + def withMetricUsageHints(hints: Set[MetricUsageHint]): MetricBuilder = { + this.copy(metricUsageHints = hints) + } + + /** + * Generates a CounterType MetricBuilder which can be used to create a counter in a StatsReceiver. + * Used to test that builder class correctly propagates configured metadata. + * @return a Countertype MetricBuilder. + */ + private[MetricBuilder] final def counterSchema: MetricBuilder = + this.copy(metricType = CounterType) + + /** + * Generates a GaugeType MetricBuilder which can be used to create a gauge in a StatsReceiver. + * Used to test that builder class correctly propagates configured metadata. + * @return a GaugeType MetricBuilder. + */ + private[MetricBuilder] final def gaugeSchema: MetricBuilder = this.copy(metricType = GaugeType) + + /** + * Generates a HistogramType MetricBuilder which can be used to create a histogram in a StatsReceiver. + * Used to test that builder class correctly propagates configured metadata. + * @return a HistogramType MetricBuilder. + */ + private[MetricBuilder] final def histogramSchema: MetricBuilder = + this.copy(metricType = HistogramType) + + /** + * Generates a CounterishGaugeType MetricBuilder which can be used to create a counter-ish gauge + * in a StatsReceiver. Used to test that builder class correctly propagates configured metadata. + * @return a CounterishGaugeType MetricBuilder. + */ + private[MetricBuilder] final def counterishGaugeSchema: MetricBuilder = + this.copy(metricType = CounterishGaugeType) + + def canEqual(other: Any): Boolean = other.isInstanceOf[MetricBuilder] + + override def equals(other: Any): Boolean = other match { + case that: MetricBuilder => + that.canEqual(this) && + keyIndicator == that.keyIndicator && + description == that.description && + units == that.units && + role == that.role && + verbosity == that.verbosity && + sourceClass == that.sourceClass && + identity == that.identity && + relativeName == that.relativeName && + processPath == that.processPath && + histogramFormat == histogramFormat && + percentiles == that.percentiles && + metricType == that.metricType + case _ => false + } + + override def hashCode(): Int = { + val state = + Array[AnyRef]( + description, + units, + role, + verbosity, + sourceClass, + identity, + relativeName, + processPath, + histogramFormat, + percentiles, + metricType) + + state.foldLeft(keyIndicator.hashCode())((a, b) => 31 * a + b.hashCode()) + } + + override def toString(): String = { + s"MetricBuilder($keyIndicator, $description, $units, $role, $verbosity, $sourceClass, $identity, $relativeName, $processPath, $histogramFormat, $percentiles, $metricType)" + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUnit.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUnit.scala new file mode 100644 index 0000000000..b52c9e5946 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUnit.scala @@ -0,0 +1,30 @@ +package com.twitter.finagle.stats + +/** + * Represents the units this metric is measured in. + * + * Common units for metrics are: + * Bytes/Kilobytes/Megabytes (for payload size, data written to disk) + * Milliseconds (for latency, GC durations) + * Requests (for successes, failures, and requests) + * Percentage (for CPU Util, Memory Usage) + */ +sealed trait MetricUnit { + + /** + * Java-friendly helper for accessing the object itself. + */ + def getInstance(): MetricUnit = this +} +case object Unspecified extends MetricUnit +case object Bytes extends MetricUnit +case object Kilobytes extends MetricUnit +case object Megabytes extends MetricUnit +case object Seconds extends MetricUnit +case object Milliseconds extends MetricUnit +case object Microseconds extends MetricUnit +case object Requests extends MetricUnit +case object Percentage extends MetricUnit +case class CustomUnit(name: String) extends MetricUnit { + override val toString: String = s"CustomUnit($name)" +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUsageHint.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUsageHint.scala new file mode 100644 index 0000000000..58b39c7759 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/MetricUsageHint.scala @@ -0,0 +1,7 @@ +package com.twitter.finagle.stats + +sealed trait MetricUsageHint {} + +object MetricUsageHint { + object HighContention extends MetricUsageHint +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/NameTranslatingStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/NameTranslatingStatsReceiver.scala index 6a25303bb9..0afedb4af1 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/NameTranslatingStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/NameTranslatingStatsReceiver.scala @@ -2,29 +2,77 @@ package com.twitter.finagle.stats /** * A StatsReceiver receiver proxy that translates all counter, stat, and gauge - * names according to a `translate` function. + * names according to a `translate(name: Seq[String])` function. * * @param self The underlying StatsReceiver to which translated names are passed - * * @param namespacePrefix the namespace used for translations */ abstract class NameTranslatingStatsReceiver( - protected val self: StatsReceiver, - namespacePrefix: String -) extends StatsReceiverProxy { + self: StatsReceiver, + namespacePrefix: String, + mode: NameTranslatingStatsReceiver.Mode) + extends TranslatingStatsReceiver(self) { + + import NameTranslatingStatsReceiver._ + + def this(self: StatsReceiver, namespacePrefix: String) = + this(self, namespacePrefix, NameTranslatingStatsReceiver.HierarchicalOnly) def this(self: StatsReceiver) = this(self, "") protected def translate(name: Seq[String]): Seq[String] - override def counter(verbosity: Verbosity, name: String*): Counter = - self.counter(verbosity, translate(name): _*) + final protected def translate(builder: MetricBuilder): MetricBuilder = { + val nextIdentity = mode match { + case FullTranslation => + builder.identity.copy( + hierarchicalName = translate(builder.identity.hierarchicalName), + dimensionalName = translate(builder.identity.dimensionalName)) + + case HierarchicalOnly => + builder.identity.copy(hierarchicalName = translate(builder.identity.hierarchicalName)) + + case DimensionalOnly => + builder.identity.copy(dimensionalName = translate(builder.identity.dimensionalName)) + } + + builder.withIdentity(nextIdentity) + } + + override def toString: String = { + mode match { + case FullTranslation | HierarchicalOnly => s"$self/$namespacePrefix" + case DimensionalOnly => self.toString + } + } +} + +object NameTranslatingStatsReceiver { + + /** The mode with which to translate stats, either hierarchical, dimensional, or both. */ + sealed trait Mode + + /** Translate metric names for both hierarchical and dimensional representations */ + case object FullTranslation extends Mode - override def stat(verbosity: Verbosity, name: String*): Stat = - self.stat(verbosity, translate(name): _*) + /** Translate metric names only for hierarchical representations */ + case object HierarchicalOnly extends Mode - override def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge = - self.addGauge(verbosity, translate(name): _*)(f) + /** Translate metric names only for dimensional representations */ + case object DimensionalOnly extends Mode - override def toString: String = s"$self/$namespacePrefix" + /** + * Translate the scope of a [[StatsReceiver]] + * + * @param parent the [[StatsReceiver]] to delegate the translated name to. + * @param namespace the segment to prepend to the name using the appropriate scope delimiter. + * @param mode select which representations to translate. + * */ + final class ScopeTranslatingStatsReceiver( + parent: StatsReceiver, + namespace: String, + mode: NameTranslatingStatsReceiver.Mode) + extends NameTranslatingStatsReceiver(parent, namespace, mode) { + protected def translate(name: Seq[String]): Seq[String] = namespace +: name + } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/NullStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/NullStatsReceiver.scala index 1a02801c61..0926044b20 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/NullStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/NullStatsReceiver.scala @@ -1,7 +1,22 @@ package com.twitter.finagle.stats object NullStatsReceiver extends NullStatsReceiver { - def get() = this + def get(): NullStatsReceiver.type = this + + val NullCounter: Counter = new Counter { + def incr(delta: Long): Unit = () + def metadata: Metadata = NoMetadata + } + + val NullStat: Stat = new Stat { + def add(value: Float): Unit = () + def metadata: Metadata = NoMetadata + } + + val NullGauge: Gauge = new Gauge { + def remove(): Unit = () + def metadata: Metadata = NoMetadata + } } /** @@ -10,17 +25,13 @@ object NullStatsReceiver extends NullStatsReceiver { * required. */ class NullStatsReceiver extends StatsReceiver { - def repr: NullStatsReceiver = this - - private[this] val NullCounter = new Counter { def incr(delta: Long): Unit = () } - private[this] val NullStat = new Stat { def add(value: Float): Unit = () } - private[this] val NullGauge = new Gauge { def remove(): Unit = () } + import NullStatsReceiver._ - def counter(verbosity: Verbosity, name: String*): Counter = NullCounter - def stat(verbosity: Verbosity, name: String*): Stat = NullStat - def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge = NullGauge + def repr: NullStatsReceiver = this - override def provideGauge(name: String*)(f: => Float): Unit = () + def counter(metricBuilder: MetricBuilder) = NullCounter + def stat(metricBuilder: MetricBuilder) = NullStat + def addGauge(metricBuilder: MetricBuilder)(f: => Float) = NullGauge override def scope(namespace: String): StatsReceiver = this @@ -29,4 +40,5 @@ class NullStatsReceiver extends StatsReceiver { override def isNull: Boolean = true override def toString: String = "NullStatsReceiver" + } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/RelativeNameMarkingStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/RelativeNameMarkingStatsReceiver.scala new file mode 100644 index 0000000000..0f7bbe2102 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/RelativeNameMarkingStatsReceiver.scala @@ -0,0 +1,28 @@ +package com.twitter.finagle.stats + +/** + * A StatsReceiver receiver proxy that populates relativeName for all counter, stat, and gauge + * schemas/metadata. + * + * @param self The underlying StatsReceiver to which modified schemas are passed + */ +case class RelativeNameMarkingStatsReceiver( + protected val self: StatsReceiver) + extends StatsReceiverProxy { + + override def counter(metricBuilder: MetricBuilder): Counter = { + val schema = metricBuilder.withRelativeName(metricBuilder.name: _*) + self.counter(schema) + } + + override def stat(metricBuilder: MetricBuilder): Stat = { + val schema = metricBuilder.withRelativeName(metricBuilder.name: _*) + self.stat(schema) + } + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + val schema = metricBuilder + .withRelativeName(metricBuilder.name: _*) + self.addGauge(schema)(f) + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiver.scala new file mode 100644 index 0000000000..1c9e7bb960 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiver.scala @@ -0,0 +1,38 @@ +package com.twitter.finagle.stats + +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.util.Try + +/** + * A StatsReceiver proxy that configures all counter, stat, and gauge + * SourceRoles to the passed in "role". + * + * @param self The underlying StatsReceiver to which translated names are passed + * @param role the role used for SourceRole Metadata + */ +case class RoleConfiguredStatsReceiver( + protected val self: StatsReceiver, + role: SourceRole, + name: Option[String] = None) + extends StatsReceiverProxy { + + override def counter(metricBuilder: MetricBuilder): Counter = { + self.counter(metricBuilder.withRole(role)) + } + + override def stat(metricBuilder: MetricBuilder): Stat = { + self.stat(metricBuilder.withRole(role)) + } + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + self.addGauge(metricBuilder.withRole(role))(f) + } + + override def registerExpression(expressionSchema: ExpressionSchema): Try[Unit] = { + val configuredExprSchema = name match { + case Some(serviceName) => expressionSchema.withRole(role).withServiceName(serviceName) + case None => expressionSchema.withRole(role) + } + self.registerExpression(configuredExprSchema) + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/RollupStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/RollupStatsReceiver.scala index 2e030495f1..0b16ce0c6e 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/RollupStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/RollupStatsReceiver.scala @@ -9,37 +9,46 @@ package com.twitter.finagle.stats * - "/errors" * - "/errors/clientErrors" * - "/errors/clientErrors/java_net_ConnectException" + * + * @param self the [[StatsReceiver]] to proxy metrics creation to + * @param hierarchicalOnly whether non-root scopes should be created with the HierarchicalOnly metrics identity */ -class RollupStatsReceiver(protected val self: StatsReceiver) extends StatsReceiverProxy { +class RollupStatsReceiver(protected val self: StatsReceiver, hierarchicalOnly: Boolean = false) + extends StatsReceiverProxy { - private[this] def tails[A](s: Seq[A]): Seq[Seq[A]] = { + /** + * @return Seq(metrics namespace -> whether the namespace is the "parent" scope) + */ + private[this] def tails[A](s: Seq[A]): Seq[(Seq[A], Boolean)] = { s match { case s @ Seq(_) => - Seq(s) + Seq(s -> true) case Seq(hd, tl @ _*) => - Seq(Seq(hd)) ++ (tails(tl) map { t => - Seq(hd) ++ t - }) + Seq(Seq(hd) -> true) ++ (tails(tl) map { case (t, _) => (Seq(hd) ++ t) -> false }) } } - override def counter(verbosity: Verbosity, names: String*): Counter = new Counter { - private[this] val allCounters = BroadcastCounter( - tails(names).map(n => self.counter(verbosity, n: _*)) - ) - def incr(delta: Long): Unit = allCounters.incr(delta) + private[this] def metrics(parent: MetricBuilder): Seq[MetricBuilder] = tails(parent.name).map { + case (name, isRoot) => + val builder = parent.withName(name: _*) + if (isRoot || !hierarchicalOnly) builder else builder.withHierarchicalOnly } - override def stat(verbosity: Verbosity, names: String*): Stat = new Stat { - private[this] val allStats = BroadcastStat( - tails(names).map(n => self.stat(verbosity, n: _*)) - ) + override def counter(metricBuilder: MetricBuilder): Counter = new Counter { + private[this] val allCounters = BroadcastCounter(metrics(metricBuilder).map(self.counter)) + def incr(delta: Long): Unit = allCounters.incr(delta) + def metadata: Metadata = allCounters.metadata + } + override def stat(metricBuilder: MetricBuilder): Stat = new Stat { + private[this] val allStats = BroadcastStat(metrics(metricBuilder).map(self.stat)) def add(value: Float): Unit = allStats.add(value) + def metadata: Metadata = allStats.metadata } - override def addGauge(verbosity: Verbosity, names: String*)(f: => Float): Gauge = new Gauge { - private[this] val underlying = tails(names).map(n => self.addGauge(verbosity, n: _*)(f)) + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = new Gauge { + private[this] val underlying = metrics(metricBuilder).map { self.addGauge(_)(f) } def remove(): Unit = underlying.foreach(_.remove()) + def metadata: Metadata = MultiMetadata(underlying.map(_.metadata)) } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/SchemaRegistry.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/SchemaRegistry.scala new file mode 100644 index 0000000000..df63b1e2c7 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/SchemaRegistry.scala @@ -0,0 +1,19 @@ +package com.twitter.finagle.stats + +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.finagle.stats.exp.ExpressionSchemaKey + +/** + * Interface used via the LoadService mechanism to obtain an + * efficient mechanism to sample stats. + */ +private[twitter] trait SchemaRegistry { + + /** Whether or not the counters are latched. */ + def hasLatchedCounters: Boolean + + def schemas(): Map[String, MetricBuilder] + + def expressions(): Map[ExpressionSchemaKey, ExpressionSchema] + +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala index 1e358e5182..a46289f8be 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/Stat.scala @@ -8,12 +8,18 @@ import scala.util.control.NonFatal * An append-only collection of time-series data. Example Stats are * "queue depth" or "query width in a stream of requests". * + * @define StatTimeScaladocLink + * [[com.twitter.finagle.stats.Stat.time[A](stat:com\.twitter\.finagle\.stats\.Stat,unit:java\.util\.concurrent\.TimeUnit)(f:=>A):A* time(Stat,TimeUnit)]] + * + * @define StatTimeFutureScaladocLink + * [[com.twitter.finagle.stats.Stat.timeFuture[A](stat:com\.twitter\.finagle\.stats\.Stat,unit:java\.util\.concurrent\.TimeUnit)(f:=>com\.twitter\.util\.Future[A]):com\.twitter\.util\.Future[A]* timeFuture(Stat,TimeUnit)]] + * * Utilities for timing synchronous execution and asynchronous - * execution are on the companion object ([[Stat.time(Stat)]] and - * [[Stat.timeFuture(Stat)]]. + * execution are on the companion object ($StatTimeScaladocLink and $StatTimeFutureScaladocLink). */ trait Stat { def add(value: Float): Unit + def metadata: Metadata } /** @@ -31,7 +37,7 @@ object Stat { try { f } finally { - stat.add(elapsed().inUnit(unit)) + stat.add(elapsed().inUnit(unit).toFloat) } } @@ -47,11 +53,11 @@ object Stat { val start = Stopwatch.timeNanos() try { f.respond { _ => - stat.add(unit.convert(Stopwatch.timeNanos() - start, TimeUnit.NANOSECONDS)) + stat.add(unit.convert(Stopwatch.timeNanos() - start, TimeUnit.NANOSECONDS).toFloat) } } catch { case NonFatal(e) => - stat.add(unit.convert(Stopwatch.timeNanos() - start, TimeUnit.NANOSECONDS)) + stat.add(unit.convert(Stopwatch.timeNanos() - start, TimeUnit.NANOSECONDS).toFloat) Future.exception(e) } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiver.scala index d2461a74a0..e984e7bec5 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiver.scala @@ -1,7 +1,19 @@ package com.twitter.finagle.stats +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.CounterishGaugeType +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.MetricBuilder.MetricType +import com.twitter.finagle.stats.MetricBuilder.UnlatchedCounter +import com.twitter.finagle.stats.NameTranslatingStatsReceiver.ScopeTranslatingStatsReceiver +import com.twitter.finagle.stats.TranslatingStatsReceiver.LabelTranslatingStatsReceiver +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.util.Return +import com.twitter.util.Try import java.lang.{Float => JFloat} import java.util.concurrent.Callable +import java.util.function.Supplier import scala.annotation.varargs /** @@ -14,12 +26,21 @@ object Verbosity { /** * Indicates that a given metric is for standard operations. */ - val Default: Verbosity = new Verbosity("Verbosity(default)") + val Default: Verbosity = new Verbosity("default") /** * Indicates that a given metric may only be useful for a debugging/troubleshooting purposes. */ - val Debug: Verbosity = new Verbosity("Verbosity(debug)") + val Debug: Verbosity = new Verbosity("debug") + + /** + * Indicates that the metrics value is short lived and the backend metric store is free to + * drop it more eagerly than normal metrics. + * + * @note there is no guarantee as to how or if the metrics store will honor the designation + * and if not implementations should interpret it the same as `Default`. + */ + val ShortLived: Verbosity = new Verbosity("short_lived") } object StatsReceiver { @@ -28,37 +49,29 @@ object StatsReceiver { /** * [[StatsReceiver]] utility methods for ease of use from java. + * + * @define ProvideGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.provideGauge(name:String*)(f:=>Float):Unit* provideGauge(String*)(=>Float)]] */ +@deprecated("Use StatsReceiver addGauge and provideGauge methods directly", "2020-06-10") object StatsReceivers { /** - * Java compatible version of [[StatsReceiver.counter]]. - */ - @varargs - @deprecated("Use vararg StatsReceiver.counter() instead", "2017-6-16") - def counter(statsReceiver: StatsReceiver, name: String*): Counter = - statsReceiver.counter(name: _*) - - /** - * Java compatible version of [[StatsReceiver.addGauge]]. + * Java compatible version of [[addGauge]]. */ @varargs - def addGauge(statsReceiver: StatsReceiver, callable: Callable[JFloat], name: String*): Gauge = + def addGauge(statsReceiver: StatsReceiver, callable: Callable[JFloat], name: String*): Gauge = { + // scalafix:off StoreGaugesAsMemberVariables statsReceiver.addGauge(name: _*)(callable.call()) + // scalafix:on StoreGaugesAsMemberVariables + } /** - * Java compatible version of [[StatsReceiver.provideGauge]]. + * Java compatible version of $ProvideGaugeScaladocLink. */ @varargs def provideGauge(statsReceiver: StatsReceiver, callable: Callable[JFloat], name: String*): Unit = statsReceiver.provideGauge(name: _*)(callable.call()) - - /** - * Java compatible version of [[StatsReceiver.stat]]. - */ - @varargs - @deprecated("Use vararg StatsReceiver.stat() instead", "2017-6-16") - def stat(statsReceiver: StatsReceiver, name: String*): Stat = statsReceiver.stat(name: _*) } /** @@ -77,7 +90,23 @@ object StatsReceivers { * Metrics created w/o an explicitly specified [[Verbosity]] level, will use [[Verbosity.Default]]. * Use [[VerbosityAdjustingStatsReceiver]] to adjust this behaviour. * - * @see [[StatsReceivers]] for a Java-friendly API. + * @define ProvideGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.provideGauge(name:String*)(f:=>Float):Unit* provideGauge(String*)(=>Float)]] + * + * @define AddGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.addGauge(name:String*)(f:=>Float):com\.twitter\.finagle\.stats\.Gauge* addGauge(String*)(=>Float)]] + * + * @define JavaProvideGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.provideGauge(f:java\.util\.function\.Supplier[Float],name:String*):Unit* provideGauge(Supplier[Float],String*)]] + * + * @define JavaAddGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.addGauge(f:java\.util\.function\.Supplier[Float],name:String*):com\.twitter\.finagle\.stats\.Gauge* addGauge(Supplier[Float],String*)]] + * + * @define AddVerboseGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.addGauge(verbosity:com\.twitter\.finagle\.stats\.Verbosity,name:String*)(f:=>Float):com\.twitter\.finagle\.stats\.Gauge* addGauge(Verbosity,String*)(=>Float)]] + * + * @define JavaVerboseAddGaugeScaladocLink + * [[com.twitter.finagle.stats.StatsReceiver.addGauge(f:java\.util\.function\.Supplier[Float],verbosity:com\.twitter\.finagle\.stats\.Verbosity,name:String*):com\.twitter\.finagle\.stats\.Gauge* addGauge(Supplier[Float],Verbosity,String*)]] */ trait StatsReceiver { @@ -95,6 +124,13 @@ trait StatsReceiver { */ def isNull: Boolean = false + /** + * Get a [[MetricBuilder metricBuilder]] for this StatsReceiver. + */ + @deprecated("Construct a MetricBuilder using its apply method", "2022-05-11") + def metricBuilder(metricType: MetricType): MetricBuilder = + MetricBuilder(metricType = metricType) + /** * Get a [[Counter counter]] with the given `name`. */ @@ -102,19 +138,43 @@ trait StatsReceiver { def counter(name: String*): Counter = counter(Verbosity.Default, name: _*) /** - * Get a [[Counter counter]] with the given `name`. + * Get a [[Counter counter]] with the given `description` and `name`. */ @varargs - def counter(verbosity: Verbosity, name: String*): Counter + final def counter(description: Some[String], name: String*): Counter = + counter(description.get, Verbosity.Default, name: _*) /** * Get a [[Counter counter]] with the given `name`. - * - * This method is a convenience for Java programs, but is no longer needed because - * [[StatsReceivers.counter]] is usable from java. */ - @deprecated("Use vararg counter() instead", "2017-6-16") - def counter0(name: String): Counter = counter(name) + @varargs + final def counter(verbosity: Verbosity, name: String*): Counter = { + val builder = this + .metricBuilder(CounterType) + .withVerbosity(verbosity) + .withName(name: _*) + + counter(builder) + } + + /** + * Get a [[Counter counter]] with the given `description` and `name`. + */ + @varargs + final def counter(description: String, verbosity: Verbosity, name: String*): Counter = { + val builder = this + .metricBuilder(CounterType) + .withVerbosity(verbosity) + .withName(name: _*) + .withDescription(description) + + counter(builder) + } + + /** + * Get a [[Counter counter]] with the given schema. + */ + def counter(schema: MetricBuilder): Counter /** * Get a [[Stat stat]] with the given name. @@ -122,19 +182,44 @@ trait StatsReceiver { @varargs def stat(name: String*): Stat = stat(Verbosity.Default, name: _*) + /** + * Get a [[Stat stat]] with the given `description` and `name`. + */ + @varargs + final def stat(description: Some[String], name: String*): Stat = + stat(description.get, Verbosity.Default, name: _*) + /** * Get a [[Stat stat]] with the given name. */ @varargs - def stat(verbosity: Verbosity, name: String*): Stat + final def stat(verbosity: Verbosity, name: String*): Stat = { + val builder = this + .metricBuilder(HistogramType) + .withVerbosity(verbosity) + .withName(name: _*) + + stat(builder) + } + + /** + * Get a [[Stat stat]] with the given `description` and `name`. + */ + @varargs + final def stat(description: String, verbosity: Verbosity, name: String*): Stat = { + val builder = this + .metricBuilder(HistogramType) + .withVerbosity(verbosity) + .withName(name: _*) + .withDescription(description) + + stat(builder) + } /** - * Get a [[Stat stat]] with the given name. This method is a convenience for Java - * programs, but is no longer needed because [[StatsReceivers.counter]] is - * usable from java. + * Get a [[Stat stat]] with the given schema. */ - @deprecated("Use vararg stat() instead", "2017-6-16") - def stat0(name: String): Stat = stat(name) + def stat(schema: MetricBuilder): Stat /** * Register a function `f` as a [[Gauge gauge]] with the given name that has @@ -144,8 +229,10 @@ trait StatsReceiver { * * Measurements under the same name are added together. * - * @see [[StatsReceiver.addGauge]] if you can properly control the lifecycle + * @see $AddGaugeScaladocLink if you can properly control the lifecycle * of the returned [[Gauge gauge]]. + * + * @see $JavaProvideGaugeScaladocLink for a Java-friendly version. */ def provideGauge(name: String*)(f: => Float): Unit = { val gauge = addGauge(name: _*)(f) @@ -154,6 +241,13 @@ trait StatsReceiver { } } + /** + * Just like $ProvideGaugeScaladocLink but optimized for better Java experience. + */ + @varargs + def provideGauge(f: Supplier[Float], name: String*): Unit = + provideGauge(name: _*)(f.get()) + /** * Add the function `f` as a [[Gauge gauge]] with the given name. * @@ -166,13 +260,21 @@ trait StatsReceiver { * * Measurements under the same name are added together. * - * @see [[StatsReceiver.provideGauge]] when there is not a good location + * @see $ProvideGaugeScaladocLink when there is not a good location * to store the returned [[Gauge gauge]] that can give the desired lifecycle. * - * @see [[http://docs.oracle.com/javase/7/docs/api/java/lang/ref/WeakReference.html java.lang.ref.WeakReference]] + * @see $JavaAddGaugeScaladocLink for a Java-friendly version. + * + * @see [[https://docs.oracle.com/javase/7/docs/api/java/lang/ref/WeakReference.html java.lang.ref.WeakReference]] */ def addGauge(name: String*)(f: => Float): Gauge = addGauge(Verbosity.Default, name: _*)(f) + /** + * Add the function `f` as a [[Gauge gauge]] with the given name and description. + */ + def addGauge(description: Some[String], name: String*)(f: => Float): Gauge = + addGauge(description.get, Verbosity.Default, name: _*)(f) + /** * Add the function `f` as a [[Gauge gauge]] with the given name. * @@ -185,12 +287,80 @@ trait StatsReceiver { * * Measurements under the same name are added together. * - * @see [[StatsReceiver.provideGauge]] when there is not a good location + * @see $ProvideGaugeScaladocLink when there is not a good location * to store the returned [[Gauge gauge]] that can give the desired lifecycle. * - * @see [[http://docs.oracle.com/javase/7/docs/api/java/lang/ref/WeakReference.html java.lang.ref.WeakReference]] + * @see $JavaVerboseAddGaugeScaladocLink for a Java-friendly version. + * + * @see [[https://docs.oracle.com/javase/7/docs/api/java/lang/ref/WeakReference.html java.lang.ref.WeakReference]] + */ + final def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge = { + val builder = this + .metricBuilder(GaugeType) + .withVerbosity(verbosity) + .withName(name: _*) + + addGauge(builder)(f) + } + + /** + * Add the function `f` as a [[Gauge gauge]] with the given name and description. */ - def addGauge(verbosity: Verbosity, name: String*)(f: => Float): Gauge + final def addGauge( + description: String, + verbosity: Verbosity, + name: String* + )( + f: => Float + ): Gauge = { + val builder = this + .metricBuilder(GaugeType) + .withVerbosity(verbosity) + .withName(name: _*) + .withDescription(description) + + addGauge(builder)(f) + } + + /** + * Just like $AddGaugeScaladocLink but optimized for better Java experience. + */ + @varargs + def addGauge(f: Supplier[Float], name: String*): Gauge = addGauge(f, Verbosity.Default, name: _*) + + /** + * Just like $AddVerboseGaugeScaladocLink but optimized for better Java experience. + */ + @varargs + final def addGauge(f: Supplier[Float], verbosity: Verbosity, name: String*): Gauge = { + val builder = this + .metricBuilder(GaugeType) + .withVerbosity(verbosity) + .withName(name: _*) + + addGauge( + if (name.length > 1) builder.withHierarchicalOnly + else builder.withDimensionalSupport + )(f.get()) + } + + /** + * Add the function `f` as a [[Gauge gauge]] with the given name. + * + * The returned [[Gauge gauge]] value is only weakly referenced by the + * [[StatsReceiver]], and if garbage collected will eventually cease to + * be a part of this measurement: thus, it needs to be retained by the + * caller. Or put another way, the measurement is only guaranteed to exist + * as long as there exists a strong reference to the returned + * [[Gauge gauge]] and typically should be stored in a member variable. + * + * Measurements under the same name are added together. + * + * @see $ProvideGaugeScaladocLink when there is not a + * good location to store the returned [[Gauge gauge]] that can give the desired lifecycle. + * @see [[https://docs.oracle.com/javase/7/docs/api/java/lang/ref/WeakReference.html java.lang.ref.WeakReference]] + */ + def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge /** * Prepend `namespace` to the names of the returned [[StatsReceiver]]. @@ -203,14 +373,63 @@ trait StatsReceiver { * }}} * * will generate [[Counter counters]] named `/client/adds` and `/client/backend/adds`. + * + * Note it's recommended to be mindful with usage of the `scope` method as it's almost always + * more efficient to pass a full metric name directly to a constructing method. + * + * Put this way, whenever possible prefer + * + * {{{ + * statsReceiver.counter("client", "adds") + * }}} + * + * to + * + * {{{ + * statsReceiver.scope("client").counter("adds") + * }}} */ def scope(namespace: String): StatsReceiver = { if (namespace == "") this - else { - new NameTranslatingStatsReceiver(this, namespace) { - protected def translate(name: Seq[String]): Seq[String] = namespace +: name - } - } + else new ScopeTranslatingStatsReceiver(this, namespace, scopeTranslation) + } + + private[finagle] def scopeTranslation: NameTranslatingStatsReceiver.Mode = + NameTranslatingStatsReceiver.HierarchicalOnly + + /** + * Create a new `StatsReceiver` that will add a scope that is only used when the metric is + * emitted in hierarchical form. + */ + def hierarchicalScope(namespace: String): StatsReceiver = { + if (namespace == "") this + else + new ScopeTranslatingStatsReceiver( + this, + namespace, + NameTranslatingStatsReceiver.HierarchicalOnly) + } + + /** + * Create a new `StatsReceiver` that will add a scope that is only used when the metric is + * emitted in dimensional form. + */ + def dimensionalScope(namespace: String): StatsReceiver = { + if (namespace == "") this + else + new ScopeTranslatingStatsReceiver( + this, + namespace, + NameTranslatingStatsReceiver.DimensionalOnly) + } + + /** + * Create a new `StatsReceiver` that will add the specified label to all created metrics. + */ + def label(labelName: String, labelValue: String): StatsReceiver = { + require(labelName.nonEmpty) + if (labelValue == "") this + else new LabelTranslatingStatsReceiver(this, labelName, labelValue) } /** @@ -223,11 +442,34 @@ trait StatsReceiver { * }}} * * will generate a [[Counter counter]] named `/client/backend/pool/adds`. + * + * Note it's recommended to be mindful with usage of the `scope` method as it's almost always + * more efficient to pass a full metric name directly to a constructing method. + * + * Put this way, whenever possible prefer + * + * {{{ + * statsReceiver.counter("client", "backend", "pool", "adds") + * }}} + * + * to + * + * {{{ + * statsReceiver.scope("client", "backend", "pool").counter("adds") + * }}} */ @varargs final def scope(namespaces: String*): StatsReceiver = namespaces.foldLeft(this)((statsReceiver, name) => statsReceiver.scope(name)) + /** + * Register an `ExpressionSchema`. + * + * Implementations that support expressions should override this to consume expressions. + */ + def registerExpression(expressionSchema: ExpressionSchema): Try[Unit] = + Return.Unit + /** * Prepend a suffix value to the next scope. * @@ -248,6 +490,47 @@ trait StatsReceiver { override def scope(namespace: String): StatsReceiver = self.scope(namespace).scope(suffix) } } + + // These two are needed to accommodate Zookeper's circular dependency. + // We'll remove them once zk is upgraded: CSL-4710. + private[stats] def counter0(name: String): Counter = counter(name) + private[stats] def stat0(name: String): Stat = stat(name) + + private[stats] def validateMetricType( + metricBuilder: MetricBuilder, + expectedType: MetricType + ): Unit = { + val typeMatch = expectedType match { + case CounterType => + metricBuilder.metricType == CounterType || metricBuilder.metricType == UnlatchedCounter + case GaugeType => + metricBuilder.metricType == GaugeType || metricBuilder.metricType == CounterishGaugeType + case _ => metricBuilder.metricType == expectedType + } + require( + typeMatch, + s"creating a $expectedType using wrong MetricBuilder: ${metricBuilder.metricType}") + } } -abstract class AbstractStatsReceiver extends StatsReceiver +abstract class AbstractStatsReceiver extends StatsReceiver { + + final def counter(metricBuilder: MetricBuilder): Counter = + counterImpl(metricBuilder.verbosity, metricBuilder.name) + + protected def counterImpl(verbosity: Verbosity, name: scala.collection.Seq[String]): Counter + + final def stat(metricBuilder: MetricBuilder): Stat = + statImpl(metricBuilder.verbosity, metricBuilder.name) + + protected def statImpl(verbosity: Verbosity, name: scala.collection.Seq[String]): Stat + + final def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = + addGaugeImpl(metricBuilder.verbosity, metricBuilder.name, f) + + protected def addGaugeImpl( + verbosity: Verbosity, + name: scala.collection.Seq[String], + f: => Float + ): Gauge +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiverProxy.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiverProxy.scala index 073233b896..3fba6d0148 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiverProxy.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsReceiverProxy.scala @@ -1,5 +1,8 @@ package com.twitter.finagle.stats +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.util.Try + /** * A proxy [[StatsReceiver]] that delegates all calls to the `self` stats receiver. */ @@ -9,10 +12,16 @@ trait StatsReceiverProxy extends StatsReceiver with DelegatingStatsReceiver { def repr: AnyRef = self.repr - def counter(verbosity: Verbosity, names: String*): Counter = self.counter(verbosity, names: _*) - def stat(verbosity: Verbosity, names: String*): Stat = self.stat(verbosity, names: _*) - def addGauge(verbosity: Verbosity, names: String*)(f: => Float): Gauge = - self.addGauge(verbosity, names: _*)(f) + def counter(metricBuilder: MetricBuilder): Counter = self.counter(metricBuilder) + def stat(metricBuilder: MetricBuilder): Stat = self.stat(metricBuilder) + def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = self.addGauge(metricBuilder)(f) + + override private[finagle] def scopeTranslation: NameTranslatingStatsReceiver.Mode = + self.scopeTranslation + + override def registerExpression( + expressionSchema: ExpressionSchema + ): Try[Unit] = self.registerExpression(expressionSchema) def underlying: Seq[StatsReceiver] = Seq(self) diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsRegistry.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsRegistry.scala index b31788f29c..d0aadc223b 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/StatsRegistry.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/StatsRegistry.scala @@ -5,6 +5,10 @@ package com.twitter.finagle.stats * efficient mechanism to sample stats. */ private[twitter] trait StatsRegistry { + + /** Whether or not the counters are latched. */ + val latched: Boolean + def apply(): Map[String, StatEntry] } @@ -19,4 +23,7 @@ private[twitter] trait StatEntry { /** The instantaneous value of the entry. */ val value: Double + + /** The type of the metric backing this StatEntry ("counter", "gauge", or "histogram"). */ + val metricType: String } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/TranslatingStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/TranslatingStatsReceiver.scala new file mode 100644 index 0000000000..e01e5a25ed --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/TranslatingStatsReceiver.scala @@ -0,0 +1,64 @@ +package com.twitter.finagle.stats + +import com.twitter.finagle.stats.MetricBuilder.Identity + +/** + * A StatsReceiver receiver proxy that translates all counter, stat, and gauge + * names according to a `translate` function. + * + * @param self The underlying StatsReceiver to which translated `MetricBuilder`s are passed + */ +abstract class TranslatingStatsReceiver( + protected val self: StatsReceiver) + extends StatsReceiverProxy { + + protected def translate(builder: MetricBuilder): MetricBuilder + + override def counter(metricBuilder: MetricBuilder): Counter = + self.counter(translate(metricBuilder)) + + override def stat(metricBuilder: MetricBuilder): Stat = + self.stat(translate(metricBuilder)) + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + self.addGauge(translate(metricBuilder))(f) + } +} + +private object TranslatingStatsReceiver { + + private final class IdentityTranslatingStatsReceiver(sr: StatsReceiver, f: Identity => Identity) + extends TranslatingStatsReceiver(sr) { + protected def translate(builder: MetricBuilder): MetricBuilder = + builder.withIdentity(f(builder.identity)) + } + + /** + * A [[TranslatingStatsReceiver]] for working with both dimensional and hierarchical metrics. + * + * Translates the [[MetricBuilder]] to prepend the label value as a scope in addition to adding + * it to the labels map. + */ + final class LabelTranslatingStatsReceiver( + sr: StatsReceiver, + labelName: String, + labelValue: String) + extends TranslatingStatsReceiver(sr) { + + require(labelName.nonEmpty) + require(labelValue.nonEmpty) + + private[this] val labelPair = labelName -> labelValue + + protected def translate(builder: MetricBuilder): MetricBuilder = + builder.withIdentity(newIdentity(builder.identity)) + + private[this] def newIdentity(identity: Identity): Identity = { + identity.copy(labels = identity.labels + labelPair) + } + } + + /** Translate the identity of a metric */ + def translateIdentity(sr: StatsReceiver)(f: Identity => Identity): StatsReceiver = + new IdentityTranslatingStatsReceiver(sr, f) +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiver.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiver.scala index 888b8e0a6f..81a074a9b9 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiver.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiver.scala @@ -6,17 +6,19 @@ package com.twitter.finagle.stats */ class VerbosityAdjustingStatsReceiver( protected val self: StatsReceiver, - defaultVerbosity: Verbosity -) extends StatsReceiverProxy { + defaultVerbosity: Verbosity) + extends StatsReceiverProxy { - override def counter(verbosity: Verbosity, names: String*): Counter = - self.counter(defaultVerbosity, names: _*) + override def stat(metricBuilder: MetricBuilder): Stat = { + self.stat(metricBuilder.withVerbosity(defaultVerbosity)) + } - override def stat(verbosity: Verbosity, names: String*): Stat = - self.stat(defaultVerbosity, names: _*) + override def counter(metricBuilder: MetricBuilder): Counter = { + self.counter(metricBuilder.withVerbosity(defaultVerbosity)) + } + + override def addGauge(metricBuilder: MetricBuilder)(f: => Float): Gauge = { + self.addGauge(metricBuilder.withVerbosity(defaultVerbosity))(f) + } - override def addGauge( - verbosity: Verbosity, - names: String* - )(f: => Float): Gauge = self.addGauge(defaultVerbosity, names: _*)(f) } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/WithHistogramDetails.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/WithHistogramDetails.scala index 1315f21950..d5da2f8632 100644 --- a/util-stats/src/main/scala/com/twitter/finagle/stats/WithHistogramDetails.scala +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/WithHistogramDetails.scala @@ -22,8 +22,7 @@ private[stats] class AggregateWithHistogramDetails(underlying: Seq[WithHistogram */ def histogramDetails: Map[String, HistogramDetail] = underlying.foldLeft[Map[String, HistogramDetail]](Map.empty[String, HistogramDetail]) { - (acc, cur) => - cur.histogramDetails ++ acc + (acc, cur) => cur.histogramDetails ++ acc } } diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Bounds.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Bounds.scala new file mode 100644 index 0000000000..49eaef4655 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Bounds.scala @@ -0,0 +1,52 @@ +package com.twitter.finagle.stats.exp + +import com.fasterxml.jackson.annotation.JsonSubTypes.Type +import com.fasterxml.jackson.annotation.{JsonSubTypes, JsonTypeInfo} + +/** + * Thresholds are specified per-expression. The default is Unbounded. + */ +@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.PROPERTY, property = "kind") +@JsonSubTypes( + Array( + new Type(value = classOf[AcceptableRange], name = "range"), + new Type(value = classOf[MonotoneThresholds], name = "monotone"), + new Type(value = classOf[Unbounded], name = "unbounded") + ) +) +sealed trait Bounds + +/** + * The state is considered healthy when the expression computation is inside this range. + * Health status and expression computation are not necessary linearly correlated. + */ +case class AcceptableRange(lowerBoundInclusive: Double, upperBoundExclusive: Double) extends Bounds + +/** + * Health status and expression computation are linearly correlated, and the healthier + * direction is indicated by [[operator]]. + * + * [lowerBoundInclusive, badThreshold] bad + * [badThreshold, goodThreshold] ok + * [goodThreshold, upperBoundExclusive] good + * + * @note all ranges are lower inclusive and upper exclusive. + * + * @param operator GreaterThan or LessThan + * @param lowerBoundInclusive Optional, undefined meaning unbounded lower bound. + * @param upperBoundExclusive Optional, undefined meaning unbounded upper bound. + */ +case class MonotoneThresholds( + operator: Operator, + badThreshold: Double, + goodThreshold: Double, + lowerBoundInclusive: Option[Double] = None, + upperBoundExclusive: Option[Double] = None) + extends Bounds + +// does not use case object because it cannot provide `Class` for @JsonSubTypes +// see https://github.com/scala/scala/pull/9279 +object Unbounded { + val get = Unbounded() +} +case class Unbounded() extends Bounds diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/DefaultExpression.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/DefaultExpression.scala new file mode 100644 index 0000000000..69c321d336 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/DefaultExpression.scala @@ -0,0 +1,81 @@ +package com.twitter.finagle.stats.exp + +import com.twitter.finagle.stats.exp.ExpressionNames._ +import com.twitter.finagle.stats.Milliseconds +import com.twitter.finagle.stats.Percentage +import com.twitter.finagle.stats.Requests + +/** + * A utility object to create common [[ExpressionSchema]]s with user defined + * [[Expression]]s. + */ +private[twitter] object DefaultExpression { + + /** + * Create an [[ExpressionSchema]] to calculate success rate with a user + * defined {@code success} [[Expression]] and a {@code failure} + * [[Expression]]. + * + * @param success an [[Expression]] to measure the success metric. + * @param failure an [[Expression]] to measure the failure metric. + * @return an [[ExpressionSchema]] that calculates success rate from success + * and failure metric. The default unit is [[Percentage]], the + * default bounds is [[MonotoneThresholds]] with operator as + * [[GreaterThan]], bad threshold as 99.5, good threshold as 99.97. + * + */ + def successRate(success: Expression, failure: Expression): ExpressionSchema = + ExpressionSchema( + successRateName, + Expression(100).multiply(success.divide(success.plus(failure)))) + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.97)) + .withUnit(Percentage) + .withDescription(s"Default $successRateName expression") + + /** + * Create an [[ExpressionSchema]] that calculates throughput with a user + * defined {@code throughput} [[Expression]]. + * + * @param throughput an [[Expression]] to measure the throughput metric + * @return an [[ExpressionSchema]] that calculates the throughput. The + * default unit is [[Requests]]. + */ + def throughput(throughput: Expression): ExpressionSchema = + requestUnitExpression(throughputName, throughput) + + /** + * Create an [[ExpressionSchema]] that calculates failures with a user + * defined {@code failures} [[Expression]]. + * + * @param failures an [[Expression]] to measure the failures metric + * @return an [[ExpressionSchema]] that calculates the failures. The + * default unit is [[Requests]]. + */ + def failures(failures: Expression): ExpressionSchema = + requestUnitExpression(failuresName, failures) + + /** + * Create an [[ExpressionSchema]] that calculates p99 latencies with a user + * defined {@code latencyP99} [[Expression]]. + * + * @param latencyP99 an [[Expression]] to measure the p99 latencies metric. + * @return an [[ExpressionSchema]] that calculates the p99 latencies. The + * default unit is [[Milliseconds]]. + */ + def latency99(latencyP99: Expression): ExpressionSchema = + ExpressionSchema(latencyName, latencyP99) + .withUnit(Milliseconds) + .withDescription("Default P99 latency expression") + + private[this] def requestUnitExpression(name: String, expression: Expression) = + ExpressionSchema(name, expression) + .withUnit(Requests) + .withDescription(s"Default $name expression") +} + +private[twitter] object ExpressionNames { + val successRateName = "success_rate" + val throughputName = "throughput" + val latencyName = "latency" + val failuresName = "failures" +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Expression.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Expression.scala new file mode 100644 index 0000000000..1fb2fdd4c9 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Expression.scala @@ -0,0 +1,206 @@ +package com.twitter.finagle.stats.exp + +import com.twitter.finagle.stats.Metadata +import com.twitter.finagle.stats.MetricBuilder +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.StatsReceiver +import scala.annotation.varargs + +private[twitter] object Expression { + + /** + * Create a expression with a constant number in double + */ + def apply(num: Double): Expression = ConstantExpression(num.toString) + + /** + * Create a histogram expression + * @param component a [[HistogramComponent]] + */ + def apply( + metadata: Metadata, + component: HistogramComponent + ): Expression = { + metadata.toMetricBuilder match { + case Some(metricBuilder) => + require( + metricBuilder.metricType == HistogramType, + "this method is for creating histogram expression") + HistogramExpression(metricBuilder, component) + case None => + NoExpression + } + } + + /** + * Create an single expression wrapping a counter or gauge. + * @param metadata MetricBuilder instance + * @param showRollup If set true, inform the expression handler to render all tail metrics + */ + def apply(metadata: Metadata, showRollup: Boolean = false): Expression = { + metadata.toMetricBuilder match { + case Some(metricBuilder) => + require(metricBuilder.metricType != HistogramType, "provide a component for histogram") + MetricExpression(metricBuilder, showRollup) + case None => + NoExpression + } + } + + /** + * Represent an expression constructed from metrics in strings, user would need to provide + * role and service name if needed to the [[ExpressionSchema]], otherwise remain unspecified. + * + * @note [[StringExpression]] is for conveniently creating expressions without getting a hold of + * [[StatsReceiver]] or [[MetricBuilder]], if they present or easy to access, users should + * prefer creating expression via Metric instances. + * StringExpressions are not guaranteed to be real metrics, please proceed with proper + * validation and testing. + * + * @param exprs a sequence of metrics in strings + */ + def apply(exprs: Seq[String], isCounter: Boolean): Expression = StringExpression(exprs, isCounter) + + /** + * Create an expression from creating a [[com.twitter.finagle.stats.Counter]] + * with a given [[StatsReceiver]]. + * + * @param statsReceiver The [[StatsReceiver]] used by the caller to record metric. + * @param name The name to create the [[com.twitter.finagle.stats.Counter]]. + * @param showRollup When true, create an expression with all tail metrics. The + * default is false. + * @param validateMetric When true, only create an expression if the counter with name + * {@code name} has been registered with the given + * {@code statsReceiver}, otherwise return `None`. When false, + * always return an expression. The default value is false. + * @return An optional [[Expression]] carrying the counter metadata. + * Return `None` if the metric of name {@code name} is not found + * in the {@code statsReceiver}. + */ + def counter( + statsReceiver: StatsReceiver, + name: Seq[String], + showRollup: Boolean = false, + validateMetric: Boolean = false + ): Option[Expression] = + if (validateMetric) None // return None for now, will add the validation in another change + else Some(this(statsReceiver.counter(name: _*).metadata, showRollup)) + + /** + * Create an expression from creating a [[com.twitter.finagle.stats.Stat]] with + * a given [[StatsReceiver]]. + * + * @param statsReceiver The [[StatsReceiver]] used by the caller to record metric. + * @param name The name to create the [[com.twitter.finagle.stats.Stat]]. + * @param histogramComponent A [[HistogramComponent]] used to create a histogram + * expression. + * @param validateMetric When true, only create an expression if the counter with + * name {@code name} has been registered with the given + * {@code statsReceiver}, otherwise return `None`. When + * false, always return an expression. The default value is + * false. + * @return An optional [[Expression]] carrying the stat metadata. + * Return `None` if the metric of name {@code name} is not + * found in the {@code statsReceiver}. + */ + def stat( + statsReceiver: StatsReceiver, + name: Seq[String], + histogramComponent: HistogramComponent, + validateMetric: Boolean = false + ): Option[Expression] = + if (validateMetric) None // return None for now, will add the validation in another change + else Some(this(statsReceiver.stat(name: _*).metadata, histogramComponent)) + + /** + * Create a list of expressions from creating a [[com.twitter.finagle.stats.Stat]] + * for each [[HistogramComponent]] with a given [[StatsReceiver]]. + * + * @param statsReceiver The [[StatsReceiver]] used by the caller to record + * metric. + * @param name The name to create the + * [[com.twitter.finagle.stats.Stat]]. + * @param histogramComponents A list of [[HistogramComponent]] used to create a + * histogram expression. The default is + * [[DefaultExpression]]. + * @param validateMetric When true, only create an expression if the counter with + * name {@code name} has been registered with the given + * {@code statsReceiver}, otherwise return `None`. When + * false, always return an expression. The default value is + * false. + * @return A sequence of [[Expression]] carrying the stat metadata. + * The order of the returned sequence is in correspondence + * with the oder of the {@code histogramComponents}. Return + * an empty sequence if the metric of name {@code name} is + * not found in the {@code statsReceiver}. + */ + def stats( + statsReceiver: StatsReceiver, + name: Seq[String], + histogramComponents: Seq[HistogramComponent] = HistogramComponent.DefaultPercentiles, + validateMetric: Boolean = false + ): Seq[Expression] = + if (validateMetric) { + Seq.empty // return an empty list for now, will add the validation in another change + } else { + val statMetadata = statsReceiver.stat(name: _*).metadata + histogramComponents.map(this(statMetadata, _)) + } +} + +/** + * Metrics with their arithmetical(or others) calculations + */ +private[twitter] sealed trait Expression { + final def plus(other: Expression): Expression = func("plus", other) + + final def minus(other: Expression): Expression = func("minus", other) + + final def divide(other: Expression): Expression = func("divide", other) + + final def multiply(other: Expression): Expression = + func("multiply", other) + + @varargs + def func(name: String, rest: Expression*): Expression = + FunctionExpression(name, this +: rest) +} + +/** + * Represents a constant double number + */ +case class ConstantExpression private[exp] (repr: String) extends Expression + +/** + * Represents compound metrics and their arithmetical(or others) calculations + */ +case class FunctionExpression private[exp] (fnName: String, exprs: Seq[Expression]) + extends Expression { + require(exprs.size != 0, "Functions must have at least 1 argument") +} + +/** + * Represents the leaf metrics + */ +case class MetricExpression private[exp] (metricBuilder: MetricBuilder, showRollup: Boolean) + extends Expression + +/** + * Represent a histogram expression with specified component, for example the average, or a percentile + * @param component a [[HistogramComponent]] + */ +case class HistogramExpression private[exp] ( + metricBuilder: MetricBuilder, + component: HistogramComponent) + extends Expression + +/** + * Represent an expression constructed from metrics in strings + * @see Expression#apply + */ +case class StringExpression private[exp] (exprs: Seq[String], isCounter: Boolean) extends Expression + +case object NoExpression extends Expression { + @varargs + override def func(name: String, rest: Expression*): Expression = NoExpression +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/ExpressionSchema.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/ExpressionSchema.scala new file mode 100644 index 0000000000..83b1e36e7c --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/ExpressionSchema.scala @@ -0,0 +1,115 @@ +package com.twitter.finagle.stats.exp + +import com.twitter.finagle.stats.HistogramFormatter +import com.twitter.finagle.stats.MetricUnit +import com.twitter.finagle.stats.SourceRole +import com.twitter.finagle.stats.Unspecified +import com.twitter.finagle.stats.exp.HistogramComponent.Avg +import com.twitter.finagle.stats.exp.HistogramComponent.Count +import com.twitter.finagle.stats.exp.HistogramComponent.Max +import com.twitter.finagle.stats.exp.HistogramComponent.Min +import com.twitter.finagle.stats.exp.HistogramComponent.Percentile +import com.twitter.finagle.stats.exp.HistogramComponent.Sum + +/** + * ExpressionSchema is builder class that construct an expression with its metadata. + * + * @param name this is going to be an important query key when fetching expressions + * @param labels service related information + * @param namespace a list of namespaces the expression belongs to, this is usually + * used to indicate a tenant in a multi-tenancy systems or similar concepts. + * For standalone services, this should be empty. + * @param expr class representation of the expression, see [[Expression]] + * @param bounds thresholds for this expression + * @param description human-readable description of an expression's significance + * @param unit the unit associated with the metrics value (milliseconds, megabytes, requests, etc) + * @param exprQuery string representation of the expression + */ +case class ExpressionSchema private ( + name: String, + labels: Map[String, String], + expr: Expression, + namespace: Seq[String], + bounds: Bounds, + description: String, + unit: MetricUnit, + exprQuery: String) { + def withBounds(bounds: Bounds): ExpressionSchema = copy(bounds = bounds) + + def withDescription(description: String): ExpressionSchema = copy(description = description) + + def withUnit(unit: MetricUnit): ExpressionSchema = copy(unit = unit) + + /** + * Configure the expression with the given namespaces. The path can be composed of segments or + * a single string. This is for multi-tenancy system tenants or similar concepts and should + * remain unset for standalone services. + */ + def withNamespace(name: String*): ExpressionSchema = + copy(namespace = this.namespace ++ name) + + def withLabel(labelName: String, labelValue: String): ExpressionSchema = + copy(labels = labels + (labelName -> labelValue)) + + private[finagle] def withRole(role: SourceRole): ExpressionSchema = + withLabel(ExpressionSchema.Role, role.toString) + + private[finagle] def withServiceName(name: String): ExpressionSchema = + withLabel(ExpressionSchema.ServiceName, name) + + def schemaKey(): ExpressionSchemaKey = { + ExpressionSchemaKey(name, labels, namespace) + } +} + +/** + * ExpressionSchemaKey is a class that exists to serve as a key into a Map of ExpressionSchemas. + * It is simply a subset of the fields of the ExpressionSchema. Namely: + * @param name + * @param labels + * @param namespaces + */ +case class ExpressionSchemaKey( + name: String, + labels: Map[String, String], + namespaces: Seq[String]) + +// expose for testing in twitter-server +private[twitter] object ExpressionSchema { + + case class ExpressionCollisionException(msg: String) extends IllegalArgumentException(msg) + + val Role: String = "role" + val ServiceName: String = "service_name" + val ProcessPath: String = "process_path" + + def apply(name: String, expr: Expression): ExpressionSchema = { + val preLabels = expr match { + case histoExpr: HistogramExpression => histogramLabel(histoExpr) + case _ => Map.empty[String, String] + } + + ExpressionSchema( + name = name, + labels = preLabels, + namespace = Seq.empty, + expr = expr, + bounds = Unbounded.get, + description = "Unspecified", + unit = Unspecified, + exprQuery = "") + } + + private def histogramLabel(histoExpr: HistogramExpression): Map[String, String] = { + val defaultHistogramFormatter = HistogramFormatter.default + val value = histoExpr.component match { + case Percentile(percentile) => defaultHistogramFormatter.labelPercentile(percentile) + case Min => defaultHistogramFormatter.labelMin + case Max => defaultHistogramFormatter.labelMax + case Avg => defaultHistogramFormatter.labelAverage + case Sum => defaultHistogramFormatter.labelSum + case Count => defaultHistogramFormatter.labelCount + } + Map("bucket" -> value) + } +} diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/HistogramComponent.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/HistogramComponent.scala new file mode 100644 index 0000000000..89485bc1e6 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/HistogramComponent.scala @@ -0,0 +1,15 @@ +package com.twitter.finagle.stats.exp + +private[twitter] object HistogramComponent { + case object Min extends HistogramComponent + case object Max extends HistogramComponent + case object Avg extends HistogramComponent + case object Sum extends HistogramComponent + case object Count extends HistogramComponent + case class Percentile(decimal: Double) extends HistogramComponent + + val DefaultPercentiles: Seq[HistogramComponent] = + Seq(Percentile(0.50), Percentile(0.99), Percentile(0.999), Percentile(0.9999)) +} + +private[twitter] sealed trait HistogramComponent diff --git a/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Operator.scala b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Operator.scala new file mode 100644 index 0000000000..4671fb1d98 --- /dev/null +++ b/util-stats/src/main/scala/com/twitter/finagle/stats/exp/Operator.scala @@ -0,0 +1,22 @@ +package com.twitter.finagle.stats.exp + +import com.fasterxml.jackson.annotation.JsonValue + +/** + * Represent Metric Expression thresholds operator. + */ +sealed trait Operator { + def isIncreasing: Boolean + @JsonValue + def toJson: String +} + +case object GreaterThan extends Operator { + def toJson: String = ">" + val isIncreasing: Boolean = true +} + +case object LessThan extends Operator { + def toJson: String = "<" + val isIncreasing: Boolean = false +} diff --git a/util-stats/src/test/java/BUILD b/util-stats/src/test/java/BUILD index da4d41aafb..ef15c458a8 100644 --- a/util-stats/src/test/java/BUILD +++ b/util-stats/src/test/java/BUILD @@ -1,9 +1,17 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "util/util-core/src/main/scala", - "util/util-stats/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-stats/src/main/java", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", ], ) diff --git a/util-stats/src/test/java/com/twitter/compiletest/MetricBuilderCompilationTest.java b/util-stats/src/test/java/com/twitter/compiletest/MetricBuilderCompilationTest.java new file mode 100644 index 0000000000..ee0bc24b72 --- /dev/null +++ b/util-stats/src/test/java/com/twitter/compiletest/MetricBuilderCompilationTest.java @@ -0,0 +1,45 @@ +package com.twitter.compiletest; + +import org.junit.Test; + +import scala.Some; + +import com.twitter.finagle.stats.InMemoryStatsReceiver; +import com.twitter.finagle.stats.MetricBuilder; +import com.twitter.finagle.stats.MetricTypes; +import com.twitter.finagle.stats.Microseconds; +import com.twitter.finagle.stats.SourceRole; +import com.twitter.finagle.stats.StatsReceiver; +import com.twitter.finagle.stats.Verbosity; + +/** + * Java compatibility test for {@link com.twitter.finagle.stats.MetricBuilder}. + */ +public final class MetricBuilderCompilationTest { + + @Test + public void testMetricBuilderConstruction() { + StatsReceiver sr = new InMemoryStatsReceiver(); + sr.metricBuilder(MetricTypes.COUNTER_TYPE); + sr.metricBuilder(MetricTypes.UNLATCHED_COUNTER_TYPE); + sr.metricBuilder(MetricTypes.COUNTERISH_GAUGE_TYPE); + sr.metricBuilder(MetricTypes.GAUGE_TYPE); + sr.metricBuilder(MetricTypes.HISTOGRAM_TYPE); + } + + @Test + public void testWithMethods() { + StatsReceiver sr = new InMemoryStatsReceiver(); + MetricBuilder mb = sr.metricBuilder(MetricTypes.COUNTER_TYPE) + .withKeyIndicator(true) + .withDescription("my cool metric") + .withVerbosity(Verbosity.Debug()) + .withSourceClass(new Some<>("com.twitter.finagle.AwesomeClass")) + .withIdentifier(new Some<>("/p/foo/bar")) + .withUnits(Microseconds.getInstance()) + .withRole(SourceRole.noRoleSpecified()) + .withName("my", "very", "cool", "name") + .withRelativeName("cool", "name") + .withPercentiles(0.99, 0.88, 0.77); + } +} diff --git a/util-stats/src/test/java/com/twitter/finagle/stats/StatsReceiverCompilationTest.java b/util-stats/src/test/java/com/twitter/compiletest/StatsReceiverCompilationTest.java similarity index 59% rename from util-stats/src/test/java/com/twitter/finagle/stats/StatsReceiverCompilationTest.java rename to util-stats/src/test/java/com/twitter/compiletest/StatsReceiverCompilationTest.java index 9407f9709e..3e58e24632 100644 --- a/util-stats/src/test/java/com/twitter/finagle/stats/StatsReceiverCompilationTest.java +++ b/util-stats/src/test/java/com/twitter/compiletest/StatsReceiverCompilationTest.java @@ -1,5 +1,14 @@ -package com.twitter.finagle.stats; - +package com.twitter.compiletest; + +import com.twitter.finagle.stats.AbstractStatsReceiver; +import com.twitter.finagle.stats.Counter; +import com.twitter.finagle.stats.Gauge; +import com.twitter.finagle.stats.InMemoryStatsReceiver; +import com.twitter.finagle.stats.JStats; +import com.twitter.finagle.stats.NullStatsReceiver; +import com.twitter.finagle.stats.Stat; +import com.twitter.finagle.stats.StatsReceiver; +import com.twitter.finagle.stats.Verbosity; import com.twitter.util.Future; import org.junit.Test; @@ -9,10 +18,9 @@ import java.util.concurrent.TimeUnit; import scala.Function0; -import scala.collection.Seq; /** - * Java compatibility layer for {@link com.twitter.finagle.stats.StatsReceiver}. + * Java compatibility test for {@link com.twitter.finagle.stats.StatsReceiver}. */ public final class StatsReceiverCompilationTest { @@ -24,7 +32,7 @@ public void testNullStatsReceiver() { @Test public void testStatMethods() { InMemoryStatsReceiver sr = new InMemoryStatsReceiver(); - Stat stat = StatsReceivers.stat(sr, "hello", "world"); + Stat stat = sr.stat("hello", "world"); Callable callable = new Callable() { public Integer call() { return 3; @@ -47,25 +55,23 @@ public Future call() { @Test public void testGaugeMethods() { InMemoryStatsReceiver sr = new InMemoryStatsReceiver(); - Callable callable = new Callable() { - public Float call() { - return 3.0f; - } - }; - Gauge gauge = StatsReceivers.addGauge(sr, callable, "hello", "world"); - StatsReceivers.provideGauge(sr, callable, "hello", "world"); - gauge.remove(); + Gauge a = sr.addGauge(() -> 3.0f, "hello", "world"); + Gauge b = sr.addGauge(() -> 4.0f, Verbosity.Debug(), "foo", "bar"); + sr.provideGauge(() -> 3.0f, "hello", "world"); + a.remove(); + b.remove(); } @Test public void testCounterMethods() { InMemoryStatsReceiver sr = new InMemoryStatsReceiver(); sr.isNull(); - Counter counter = StatsReceivers.counter(sr, "hello", "world"); + Counter counter = sr.counter("hello", "world"); counter.incr(); counter.incr(100); } + @Test public void testScope() { InMemoryStatsReceiver sr = new InMemoryStatsReceiver(); StatsReceiver sr2 = sr.scope(); @@ -80,10 +86,10 @@ public void testAbstractStatsReceiver() { StatsReceiver sr = new AbstractStatsReceiver() { @Override - public Stat stat(Verbosity verbosity, String... name) { - return nullSr.stat(verbosity, name); + public Stat statImpl(Verbosity verbosity, scala.collection.Seq name) { + return nullSr.stat(verbosity, name.toSeq()); } - + @Override public Object repr() { return this; @@ -95,18 +101,13 @@ public boolean isNull() { } @Override - public Counter counter(Verbosity verbosity, Seq name) { - return nullSr.counter(name); - } - - @Override - public Stat stat(Verbosity verbosity, Seq name) { - return nullSr.stat(name); + public Counter counterImpl(Verbosity verbosity, scala.collection.Seq name) { + return nullSr.counter(name.toSeq()); } @Override - public Gauge addGauge(Verbosity verbosity, Seq name, Function0 f) { - return nullSr.addGauge(verbosity, name, f); + public Gauge addGaugeImpl(Verbosity verbosity, scala.collection.Seq name, Function0 f) { + return nullSr.addGauge(verbosity, name.toSeq(), f); } }; } diff --git a/util-stats/src/test/scala/BUILD b/util-stats/src/test/scala/BUILD index 77b6c92de3..5cb23a4da6 100644 --- a/util-stats/src/test/scala/BUILD +++ b/util-stats/src/test/scala/BUILD @@ -1,12 +1,21 @@ junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/mockito:mockito-core", "3rdparty/jvm/org/scalacheck", "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-stats/src/main/scala", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", ], ) diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/BlacklistStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/BlacklistStatsReceiverTest.scala deleted file mode 100644 index f854caf76a..0000000000 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/BlacklistStatsReceiverTest.scala +++ /dev/null @@ -1,66 +0,0 @@ -package com.twitter.finagle.stats - -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner - -@RunWith(classOf[JUnitRunner]) -class BlacklistStatsReceiverTest extends FunSuite { - test("BlacklistStatsReceiver blacklists properly") { - val inmemory = new InMemoryStatsReceiver() - val bsr = new BlacklistStatsReceiver(inmemory, { case _ => true }) - val ctr = bsr.counter("foo", "bar") - ctr.incr() - val gauge = bsr.addGauge("foo", "baz") { 3.0f } - val stat = bsr.stat("qux") - stat.add(3) - - assert(inmemory.counters.isEmpty) - assert(inmemory.gauges.isEmpty) - assert(inmemory.stats.isEmpty) - } - - test("BlacklistStatsReceiver allows through properly") { - val inmemory = new InMemoryStatsReceiver() - val bsr = new BlacklistStatsReceiver(inmemory, { case _ => false }) - val ctr = bsr.counter("foo", "bar") - ctr.incr() - val gauge = bsr.addGauge("foo", "baz") { 3.0f } - val stat = bsr.stat("qux") - stat.add(3) - - assert(inmemory.counters == Map(Seq("foo", "bar") -> 1)) - assert(inmemory.gauges(Seq("foo", "baz"))() == 3.0) - assert(inmemory.stats == Map(Seq("qux") -> Seq(3.0))) - } - - test("BlacklistStatsReceiver can go both ways properly") { - val inmemory = new InMemoryStatsReceiver() - val bsr = new BlacklistStatsReceiver(inmemory, { case seq => seq.length != 2 }) - val ctr = bsr.counter("foo", "bar") - ctr.incr() - val gauge = bsr.addGauge("foo", "baz") { 3.0f } - val stat = bsr.stat("qux") - stat.add(3) - - assert(inmemory.counters == Map(Seq("foo", "bar") -> 1)) - assert(inmemory.gauges(Seq("foo", "baz"))() == 3.0) - assert(inmemory.stats.isEmpty) - } - - test("BlacklistStatsReceiver works scoped") { - val inmemory = new InMemoryStatsReceiver() - val bsr = - new BlacklistStatsReceiver(inmemory, { case seq => seq == Seq("foo", "bar") }).scope("foo") - val ctr = bsr.counter("foo", "bar") - ctr.incr() - val gauge = bsr.addGauge("foo", "baz") { 3.0f } - val stat = bsr.stat("bar") - stat.add(3) - - assert(inmemory.counters == Map(Seq("foo", "foo", "bar") -> 1)) - assert(inmemory.gauges(Seq("foo", "foo", "baz"))() == 3.0) - assert(inmemory.stats.isEmpty) - } - -} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/BroadcastStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/BroadcastStatsReceiverTest.scala index 5ff400166b..b3d07e9866 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/BroadcastStatsReceiverTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/BroadcastStatsReceiverTest.scala @@ -2,12 +2,10 @@ package com.twitter.finagle.stats import com.twitter.util.{Future, Await} -import org.junit.runner.RunWith -import org.scalatest.{Matchers, FunSuite} -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers -@RunWith(classOf[JUnitRunner]) -class BroadcastStatsReceiverTest extends FunSuite with Matchers { +class BroadcastStatsReceiverTest extends AnyFunSuite with Matchers { test("counter") { val recv1 = new InMemoryStatsReceiver @@ -65,7 +63,9 @@ class BroadcastStatsReceiverTest extends FunSuite with Matchers { assert(None == recv1.gauges.get(Seq("hi"))) assert(None == recv1.gauges.get(Seq("scopeA", "hi"))) + // scalafix:off StoreGaugesAsMemberVariables val gaugeA = broadcastA.addGauge("hi") { 5f } + // scalafix:on StoreGaugesAsMemberVariables assert(5f == recv1.gauges(Seq("hi"))()) assert(5f == recv1.gauges(Seq("scopeA", "hi"))()) diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/CachedRegexTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/CachedRegexTest.scala new file mode 100644 index 0000000000..bef0e3ced9 --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/CachedRegexTest.scala @@ -0,0 +1,41 @@ +package com.twitter.finagle.stats + +import scala.jdk.CollectionConverters._ +import org.scalatest.funsuite.AnyFunSuite + +class CachedRegexTest extends AnyFunSuite { + test("Lets through things it should let through") { + val cachedRegex = new CachedRegex("foo".r) + val original: Map[String, Number] = Map("bar" -> 3) + assert(cachedRegex(original) == original) + assert(cachedRegex.regexMatchCache.asScala == Map("bar" -> false)) + } + + test("Doesn't let through things it shouldn't let through") { + val cachedRegex = new CachedRegex("foo".r) + val original: Map[String, Number] = Map("foo" -> 3) + assert(cachedRegex(original) == Map.empty) + assert(cachedRegex.regexMatchCache.asScala == Map("foo" -> true)) + } + + test("Keeps track of things that are repeated") { + val cachedRegex = new CachedRegex("foo".r) + val original: Map[String, Number] = Map("foo" -> 3) + assert(cachedRegex(original) == Map.empty) + assert(cachedRegex.regexMatchCache.asScala == Map("foo" -> true)) + + assert(cachedRegex(original) == Map.empty) + assert(cachedRegex.regexMatchCache.asScala == Map("foo" -> true)) + } + + test("Strips keys that don't reappear") { + val cachedRegex = new CachedRegex("foo".r) + val original: Map[String, Number] = Map("foo" -> 3) + assert(cachedRegex(original) == Map.empty) + assert(cachedRegex.regexMatchCache.asScala == Map("foo" -> true)) + + val updated: Map[String, Number] = Map("bar" -> 4) + assert(cachedRegex(updated) == updated) + assert(cachedRegex.regexMatchCache.asScala == Map("bar" -> false)) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/CategorizingExceptionStatsHandlerTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/CategorizingExceptionStatsHandlerTest.scala index ab5ae3c492..406450a3bc 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/CategorizingExceptionStatsHandlerTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/CategorizingExceptionStatsHandlerTest.scala @@ -1,11 +1,9 @@ package com.twitter.finagle.stats -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.compat._ -@RunWith(classOf[JUnitRunner]) -class CategorizingExceptionStatsHandlerTest extends FunSuite { +class CategorizingExceptionStatsHandlerTest extends AnyFunSuite { val categorizer = (t: Throwable) => { "clienterrors" } test("uses category, source, exception chains and rollup") { @@ -13,7 +11,7 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { val esh = new CategorizingExceptionStatsHandler( _ => Some("clienterrors"), - PartialFunction.apply(_ => Some("service")), + { case _ => Some("service") }, true ) @@ -22,9 +20,9 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted - assert(receiver.counters.filterKeys(_.contains("failures")).size == 0) + assert(receiver.counters.view.filterKeys(_.contains("failures")).size == 0) - assert(receiver.counters.filterKeys(_.contains("clienterrors")).size == 3) + assert(receiver.counters.view.filterKeys(_.contains("clienterrors")).size == 3) assert(receiver.counters(Seq("clienterrors")) == 1) assert(receiver.counters(Seq("clienterrors", classOf[RuntimeException].getName)) == 1) assert( @@ -33,7 +31,7 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { ) == 1 ) - assert(receiver.counters.filterKeys(_.contains("sourcedfailures")).size == 3) + assert(receiver.counters.view.filterKeys(_.contains("sourcedfailures")).size == 3) assert(receiver.counters(Seq("sourcedfailures", "service")) == 1) assert( receiver.counters(Seq("sourcedfailures", "service", classOf[RuntimeException].getName)) == 1 @@ -70,8 +68,8 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted - assert(receiver.counters.filterKeys(_.contains("failures")).size == 3) - assert(receiver.counters.filterKeys(_.contains("sourcedfailures")).size == 0) + assert(receiver.counters.view.filterKeys(_.contains("failures")).size == 3) + assert(receiver.counters.view.filterKeys(_.contains("sourcedfailures")).size == 0) assert( keys == Seq( @@ -93,9 +91,9 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted - assert(receiver.counters.filterKeys(_.contains("failures")).size == 0) + assert(receiver.counters.view.filterKeys(_.contains("failures")).size == 0) - assert(receiver.counters.filterKeys(_.contains("clienterrors")).size == 2) + assert(receiver.counters.view.filterKeys(_.contains("clienterrors")).size == 2) assert(receiver.counters(Seq("clienterrors")) == 1) assert( receiver.counters( @@ -103,7 +101,7 @@ class CategorizingExceptionStatsHandlerTest extends FunSuite { ) == 1 ) - assert(receiver.counters.filterKeys(_.contains("sourcedfailures")).size == 2) + assert(receiver.counters.view.filterKeys(_.contains("sourcedfailures")).size == 2) assert(receiver.counters(Seq("sourcedfailures", "service")) == 1) assert( receiver.counters( diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/CumulativeGaugeTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/CumulativeGaugeTest.scala index a9d6cf5dd8..5037c85620 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/CumulativeGaugeTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/CumulativeGaugeTest.scala @@ -2,9 +2,10 @@ package com.twitter.finagle.stats import java.util.concurrent.Executor import java.util.concurrent.atomic.AtomicInteger -import org.scalatest.FunSuite +import org.scalatest.concurrent.{Eventually, IntegrationPatience} +import org.scalatest.funsuite.AnyFunSuite -class CumulativeGaugeTest extends FunSuite { +class CumulativeGaugeTest extends AnyFunSuite with Eventually with IntegrationPatience { private[this] val sameThreadExecutor = new Executor { def execute(command: Runnable): Unit = command.run() @@ -27,13 +28,13 @@ class CumulativeGaugeTest extends FunSuite { val gauge = new TestGauge() assert(0 == gauge.numRegisters.get) - gauge.addGauge { 0.0f } + gauge.addGauge({ 0.0f }, NoMetadata) assert(1 == gauge.numRegisters.get) } test("a CumulativeGauge with size = 1 should deregister when all gauges are removed") { val gauge = new TestGauge() - val added = gauge.addGauge { 1.0f } + val added = gauge.addGauge({ 1.0f }, NoMetadata) assert(0 == gauge.numDeregisters.get) added.remove() @@ -45,7 +46,7 @@ class CumulativeGaugeTest extends FunSuite { ) { val gauge = new TestGauge() assert(0 == gauge.numDeregisters.get) - val added = gauge.addGauge { 1.0f } + val added = gauge.addGauge({ 1.0f }, NoMetadata) System.gc() gauge.cleanRefs() @@ -53,35 +54,34 @@ class CumulativeGaugeTest extends FunSuite { assert(0 == gauge.numDeregisters.get) } - if (!sys.props.contains("SKIP_FLAKY")) - test( - "a CumulativeGauge with size = 1 should deregister after a System.gc when no references are held onto, after enough gets" - ) { - val gauge = new TestGauge() - var added = gauge.addGauge { 1.0f } - assert(0 == gauge.numDeregisters.get) + test( + "a CumulativeGauge with size = 1 should deregister after a System.gc when no references are held onto, after enough gets" + ) { + val gauge = new TestGauge() + var added = gauge.addGauge({ 1.0f }, NoMetadata) + assert(0 == gauge.numDeregisters.get) - added = null - System.gc() - gauge.cleanRefs() + added = null + System.gc() + eventually { + gauge.cleanRefs() assert(gauge.getValue == 0.0f) assert(gauge.numDeregisters.get > 0) } + } test("a CumulativeGauge should sum values across all registered gauges") { val gauge = new TestGauge() - val underlying = 0.until(100).map { _ => - gauge.addGauge { 10.0f } - } + val underlying = 0.until(100).map { _ => gauge.addGauge({ 10.0f }, NoMetadata) } assert(gauge.getValue == (10.0f * 100)) } test("a CumulativeGauge should discount gauges once removed") { val gauge = new TestGauge() - val underlying = Array.fill(100) { gauge.addGauge { 10.0f } } + val underlying = Array.fill(100) { gauge.addGauge({ 10.0f }, NoMetadata) } assert(gauge.getValue == (10.0f * 100)) underlying(0).remove() assert(gauge.getValue == (10.0f * 99)) diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/DefaultStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/DefaultStatsReceiverTest.scala new file mode 100644 index 0000000000..ca25a83fe2 --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/DefaultStatsReceiverTest.scala @@ -0,0 +1,15 @@ +package com.twitter.finagle.stats + +import org.scalatest.funsuite.AnyFunSuite + +class DefaultStatsReceiverTest extends AnyFunSuite { + test("DefaultStatsReceiver initializes dimensional scope and label") { + LoadedStatsReceiver.self = new InMemoryStatsReceiver + val fooCounter = DefaultStatsReceiver.counter("foo") + val identity = fooCounter.metadata.toMetricBuilder.get.identity + + assert(identity.dimensionalName == Seq("foo")) + assert(identity.hierarchicalName == Seq("foo")) + assert(identity.labels == Map("implementation" -> "app")) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/DelegatingStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/DelegatingStatsReceiverTest.scala index 287220fc5f..530a37c2b8 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/DelegatingStatsReceiverTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/DelegatingStatsReceiverTest.scala @@ -1,13 +1,10 @@ package com.twitter.finagle.stats -import org.junit.runner.RunWith import org.scalacheck.{Arbitrary, Gen} -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.prop.GeneratorDrivenPropertyChecks +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class DelegatingStatsReceiverTest extends FunSuite with GeneratorDrivenPropertyChecks { +class DelegatingStatsReceiverTest extends AnyFunSuite with ScalaCheckDrivenPropertyChecks { class Ctx { var leafCounter = 0 @@ -17,9 +14,9 @@ class DelegatingStatsReceiverTest extends FunSuite with GeneratorDrivenPropertyC Gen.const(new InMemoryStatsReceiver()) } - private[this] def blacklistStatsReceiver(depth: Int): Gen[StatsReceiver] = + private[this] def denylistStatsReceiver(depth: Int): Gen[StatsReceiver] = if (depth > 3) inMemoryStatsReceiver - else statsReceiverTopology(depth).map(new BlacklistStatsReceiver(_, { case _ => false })) + else statsReceiverTopology(depth).map(new DenylistStatsReceiver(_, { case _ => false })) private[this] def rollupStatsReceiver(depth: Int): Gen[StatsReceiver] = if (depth > 3) inMemoryStatsReceiver @@ -58,7 +55,7 @@ class DelegatingStatsReceiverTest extends FunSuite with GeneratorDrivenPropertyC private[this] def statsReceiverTopology(depth: Int): Gen[StatsReceiver] = { val seq = Seq( inMemoryStatsReceiver, - Gen.lzy(blacklistStatsReceiver(depth + 1)), + Gen.lzy(denylistStatsReceiver(depth + 1)), Gen.lzy(scopedStatsReceiver(depth + 1)), Gen.lzy(broadcastStatsReceiver(depth + 1)), Gen.lzy(proxyStatsReceiver(depth + 1)), @@ -67,14 +64,14 @@ class DelegatingStatsReceiverTest extends FunSuite with GeneratorDrivenPropertyC Gen.oneOf(seq).flatMap(identity) } - implicit val impl = Arbitrary(Gen.delay(statsReceiverTopology(0))) + implicit val impl: Arbitrary[StatsReceiver] = Arbitrary(Gen.delay(statsReceiverTopology(0))) } test("DelegatingStatsReceiver.all collects effectively across many StatsReceivers") { val ctx = new Ctx import ctx._ - forAll { statsReceiver: StatsReceiver => + forAll { (statsReceiver: StatsReceiver) => assert(DelegatingStatsReceiver.all(statsReceiver).size == leafCounter) leafCounter = 0 } diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/DenylistStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/DenylistStatsReceiverTest.scala new file mode 100644 index 0000000000..2dd90fc64a --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/DenylistStatsReceiverTest.scala @@ -0,0 +1,100 @@ +package com.twitter.finagle.stats + +import org.scalatest.funsuite.AnyFunSuite + +class DenylistStatsReceiverTest extends AnyFunSuite { + // scalafix:off StoreGaugesAsMemberVariables + test("DenylistStatsReceiver denylists properly") { + val inmemory = new InMemoryStatsReceiver() + val bsr = new DenylistStatsReceiver(inmemory, { case _ => true }) + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("foo", "baz") { 3.0f } + val stat = bsr.stat("qux") + stat.add(3) + + assert(inmemory.counters.isEmpty) + assert(inmemory.gauges.isEmpty) + assert(inmemory.stats.isEmpty) + } + + test("DenylistStatsReceiver allows through properly") { + val inmemory = new InMemoryStatsReceiver() + val bsr = new DenylistStatsReceiver(inmemory, { case _ => false }) + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("foo", "baz") { 3.0f } + val stat = bsr.stat("qux") + stat.add(3) + + assert(inmemory.counters == Map(Seq("foo", "bar") -> 1)) + assert(inmemory.gauges(Seq("foo", "baz"))() == 3.0) + assert(inmemory.stats == Map(Seq("qux") -> Seq(3.0))) + } + + test("DenylistStatsReceiver can go both ways properly") { + val inmemory = new InMemoryStatsReceiver() + val bsr = new DenylistStatsReceiver(inmemory, { case seq => seq.length != 2 }) + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("foo", "baz") { 3.0f } + val stat = bsr.stat("qux") + stat.add(3) + + assert(inmemory.counters == Map(Seq("foo", "bar") -> 1)) + assert(inmemory.gauges(Seq("foo", "baz"))() == 3.0) + assert(inmemory.stats.isEmpty) + } + + test("DenylistStatsReceiver works scoped") { + val inmemory = new InMemoryStatsReceiver() + val bsr = + new DenylistStatsReceiver(inmemory, { case seq => seq == Seq("foo", "bar") }).scope("foo") + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("foo", "baz") { 3.0f } + val stat = bsr.stat("bar") + stat.add(3) + + assert(inmemory.counters == Map(Seq("foo", "foo", "bar") -> 1)) + assert(inmemory.gauges(Seq("foo", "foo", "baz"))() == 3.0) + assert(inmemory.stats.isEmpty) + } + + test("Denylist.orElseDenied") { + val inmemory = new InMemoryStatsReceiver() + val pf: PartialFunction[Seq[String], Boolean] = { + case s if s.contains("foo") => false + } + val bsr = DenylistStatsReceiver.orElseDenied(inmemory, pf) + + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("foo", "baz") { 3.0f } + val stat = bsr.stat("qux") + stat.add(3) + + assert(inmemory.counters.nonEmpty) + assert(inmemory.gauges.nonEmpty) + assert(inmemory.stats.isEmpty) + } + + test("Denylist.orElseAdmitted") { + val inmemory = new InMemoryStatsReceiver() + val pf: PartialFunction[Seq[String], Boolean] = { + case s if s.contains("foo") => true + } + val bsr = DenylistStatsReceiver.orElseAdmitted(inmemory, pf) + + val ctr = bsr.counter("foo", "bar") + ctr.incr() + val gauge = bsr.addGauge("baz") { 3.0f } + val stat = bsr.stat("qux") + stat.add(3) + + assert(inmemory.counters.isEmpty) + assert(inmemory.gauges.nonEmpty) + assert(inmemory.stats.nonEmpty) + } + // scalafix:on StoreGaugesAsMemberVariables +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/HistogramFormatterTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/HistogramFormatterTest.scala new file mode 100644 index 0000000000..b1f24504a7 --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/HistogramFormatterTest.scala @@ -0,0 +1,43 @@ +package com.twitter.finagle.stats + +import org.scalatest.funsuite.AnyFunSuite + +class HistogramFormatterTest extends AnyFunSuite { + + test("Ostrich - formats as expected") { + assert(HistogramFormatter.Ostrich.labelAverage == "average") + assert(HistogramFormatter.Ostrich.labelMax == "maximum") + assert(HistogramFormatter.Ostrich.labelMin == "minimum") + assert(HistogramFormatter.Ostrich.labelCount == "count") + assert(HistogramFormatter.Ostrich.labelSum == "sum") + assert(HistogramFormatter.Ostrich.labelPercentile(0.50) == "p50") + assert(HistogramFormatter.Ostrich.labelPercentile(0.99) == "p99") + assert(HistogramFormatter.Ostrich.labelPercentile(0.999) == "p999") + assert(HistogramFormatter.Ostrich.labelPercentile(0.9999) == "p9999") + } + + test("CommonMetrics - formats as expected") { + assert(HistogramFormatter.CommonsMetrics.labelAverage == "avg") + assert(HistogramFormatter.CommonsMetrics.labelMax == "max") + assert(HistogramFormatter.CommonsMetrics.labelMin == "min") + assert(HistogramFormatter.CommonsMetrics.labelCount == "count") + assert(HistogramFormatter.CommonsMetrics.labelSum == "sum") + assert(HistogramFormatter.CommonsMetrics.labelPercentile(0.50) == "p50") + assert(HistogramFormatter.CommonsMetrics.labelPercentile(0.99) == "p99") + assert(HistogramFormatter.CommonsMetrics.labelPercentile(0.999) == "p9990") + assert(HistogramFormatter.CommonsMetrics.labelPercentile(0.9999) == "p9999") + } + + test("CommonStats - formats as expected") { + assert(HistogramFormatter.CommonsStats.labelAverage == "avg") + assert(HistogramFormatter.CommonsStats.labelMax == "max") + assert(HistogramFormatter.CommonsStats.labelMin == "min") + assert(HistogramFormatter.CommonsStats.labelCount == "count") + assert(HistogramFormatter.CommonsStats.labelSum == "sum") + assert(HistogramFormatter.CommonsStats.labelPercentile(0.50) == "50_0_percentile") + assert(HistogramFormatter.CommonsStats.labelPercentile(0.99) == "99_0_percentile") + assert(HistogramFormatter.CommonsStats.labelPercentile(0.999) == "99_9_percentile") + assert(HistogramFormatter.CommonsStats.labelPercentile(0.9999) == "99_99_percentile") + } + +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/InMemoryStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/InMemoryStatsReceiverTest.scala index 73576483cf..be58e5d891 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/InMemoryStatsReceiverTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/InMemoryStatsReceiverTest.scala @@ -1,15 +1,23 @@ package com.twitter.finagle.stats -import java.io.{ByteArrayOutputStream, PrintStream} +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.exp.Expression +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.finagle.stats.exp.ExpressionSchemaKey +import com.twitter.finagle.stats.exp.HistogramComponent +import java.io.ByteArrayOutputStream +import java.io.PrintStream import java.nio.charset.StandardCharsets -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.concurrent.{Eventually, IntegrationPatience} -import org.scalatest.junit.JUnitRunner +import org.scalatest.concurrent.Eventually +import org.scalatest.concurrent.IntegrationPatience +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.parallel.immutable.ParRange -@RunWith(classOf[JUnitRunner]) -class InMemoryStatsReceiverTest extends FunSuite with Eventually with IntegrationPatience { +class InMemoryStatsReceiverTest extends AnyFunSuite with Eventually with IntegrationPatience { + // scalafix:off StoreGaugesAsMemberVariables test("clear") { val inMemoryStatsReceiver = new InMemoryStatsReceiver inMemoryStatsReceiver.counter("counter").incr() @@ -32,7 +40,7 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio test("threadsafe counter") { val inMemoryStatsReceiver = new InMemoryStatsReceiver - (1 to 50).par.foreach(_ => inMemoryStatsReceiver.counter("same").incr()) + new ParRange(1 to 50).foreach(_ => inMemoryStatsReceiver.counter("same").incr()) eventually { assert(inMemoryStatsReceiver.counter("same")() == 50) } @@ -40,7 +48,7 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio test("threadsafe stats") { val inMemoryStatsReceiver = new InMemoryStatsReceiver - (1 to 50).par.foreach(_ => inMemoryStatsReceiver.stat("same").add(1.0f)) + new ParRange(1 to 50).foreach(_ => inMemoryStatsReceiver.stat("same").add(1.0f)) eventually { assert(inMemoryStatsReceiver.stat("same")().size == 50) } @@ -57,7 +65,7 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio test("ReadableGauge.toString") { var n = 0 val stats = new InMemoryStatsReceiver() - val g = stats.addGauge("a", "b") { n } + val g = stats.addGauge("a", "b") { n.toFloat } assert("Gauge(a/b=0.0)" == g.toString) n = 11 @@ -120,6 +128,16 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio assert(stats.verbosity(Seq("baz")) == Verbosity.Debug) } + test("register expressions") { + val stats = new InMemoryStatsReceiver() + stats.registerExpression( + ExpressionSchema( + "a", + Expression(MetricBuilder(name = Seq("counter"), metricType = CounterType)))) + + assert(stats.expressions.contains(ExpressionSchemaKey("a", Map(), Nil))) + } + test("print") { val baos = new ByteArrayOutputStream() val ps = new PrintStream(baos, true, "utf-8") @@ -231,7 +249,7 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio inMemoryStatsReceiver.print(ps, includeHeaders = true) val content = new String(baos.toByteArray, StandardCharsets.UTF_8) - val parts = content.split('\n') + val parts = content.filterNot(_ == '\r').split('\n') assert(parts.length == 17) assert(parts(0) == "Counters:") @@ -294,7 +312,7 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio inMemoryStatsReceiver.print(ps, includeHeaders = true) val content = new String(baos.toByteArray, StandardCharsets.UTF_8) - val parts = content.split('\n') + val parts = content.filterNot(_ == '\r').split('\n') assert(parts.length == 12) assert(parts(0) == "Counters:") @@ -313,4 +331,121 @@ class InMemoryStatsReceiverTest extends FunSuite with Eventually with Integratio ps.close() } } + // scalafix:on StoreGaugesAsMemberVariables + + test("creating metrics stores their metadata and clear()ing should remove that metadata") { + val stats = new InMemoryStatsReceiver() + stats.addGauge("coolGauge") { 3 } + assert(stats.schemas(Seq("coolGauge")).isInstanceOf[MetricBuilder]) + stats.counter("sweetCounter") + assert(stats.schemas(Seq("sweetCounter")).isInstanceOf[MetricBuilder]) + stats.stat("radHisto") + assert(stats.schemas(Seq("radHisto")).isInstanceOf[MetricBuilder]) + + stats.clear() + assert(stats.schemas.isEmpty) + } + + test("printSchemas should print schemas") { + val stats = new InMemoryStatsReceiver() + stats.addGauge("coolGauge") { 3 } + stats.counter("sweet", "counter") + stats.scope("rad").stat("histo") + + val baos = new ByteArrayOutputStream() + val ps = new PrintStream(baos, true, "utf-8") + try { + stats.printSchemas(ps) + val content = new String(baos.toByteArray, StandardCharsets.UTF_8) + val parts = content.split('\n') + + assert(parts.length == 3) + assert( + parts( + 0) == "coolGauge MetricBuilder(false, No description provided, Unspecified, NoRoleSpecified, default, None, Identity(coolGauge, coolGauge, Map(), NonDeterminate), List(), None, default, Vector(), GaugeType)") + assert( + parts( + 1) == "rad/histo MetricBuilder(false, No description provided, Unspecified, NoRoleSpecified, default, None, Identity(rad/histo, histo, Map(), NonDeterminate), List(), None, default, Vector(), HistogramType)") + assert( + parts( + 2) == "sweet/counter MetricBuilder(false, No description provided, Unspecified, NoRoleSpecified, default, None, Identity(sweet/counter, sweet_counter, Map(), NonDeterminate), List(), None, default, Vector(), CounterType)") + } finally { + ps.close() + } + } + + test("expressions can be hydrated from metadata that metrics use") { + val sr = new InMemoryStatsReceiver + + val aCounter = sr.scope("test").counter("a") + val bHisto = sr.scope("test").stat("b") + val cGauge = sr.scope("test").addGauge("c") { 1 } + + val expression = sr.registerExpression( + ExpressionSchema( + "test_expression", + Expression(aCounter.metadata).plus(Expression(bHisto.metadata, HistogramComponent.Min) + .plus(Expression(cGauge.metadata))) + )) + + // what we expected as hydrated metric builders + val aaSchema = + MetricBuilder(metricType = CounterType) + .withIdentity( + MetricBuilder.Identity( + hierarchicalName = Seq("test", "a"), + dimensionalName = Seq("a"), + labels = Map.empty)) + + val bbSchema = + MetricBuilder(metricType = HistogramType) + .withIdentity( + MetricBuilder.Identity( + hierarchicalName = Seq("test", "b"), + dimensionalName = Seq("b"), + labels = Map.empty)) + + val ccSchema = + MetricBuilder(metricType = GaugeType) + .withIdentity( + MetricBuilder.Identity( + hierarchicalName = Seq("test", "c"), + dimensionalName = Seq("c"), + labels = Map.empty)) + + val expected_expression = ExpressionSchema( + "test_expression", + Expression(aaSchema).plus( + Expression(bbSchema, HistogramComponent.Min).plus(Expression(ccSchema)))) + + assert(sr.expressions + .get(ExpressionSchemaKey("test_expression", Map(), Nil)).get.expr == expected_expression.expr) + } + + test("getExpressionsWithLabel") { + val sr = new InMemoryStatsReceiver + + val aCounter = sr.scope("test").counter("a") + + sr.registerExpression( + ExpressionSchema( + "test_expression_0", + Expression(aCounter.metadata) + ).withLabel(ExpressionSchema.Role, "test")) + + sr.registerExpression( + ExpressionSchema( + "test_expression_1", + Expression(aCounter.metadata) + ).withLabel(ExpressionSchema.Role, "dont appear")) + + val withLabel = sr.getAllExpressionsWithLabel(ExpressionSchema.Role, "test") + assert(withLabel.size == 1) + + assert( + !withLabel.contains( + ExpressionSchemaKey("test_expression_1", Map((ExpressionSchema.Role, "dont appear")), Nil))) + + assert(sr.getAllExpressionsWithLabel("nonexistent key", "nonexistent value").isEmpty) + } } diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/LazyStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/LazyStatsReceiverTest.scala new file mode 100644 index 0000000000..67dea43e02 --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/LazyStatsReceiverTest.scala @@ -0,0 +1,23 @@ +package com.twitter.finagle.stats + +import org.scalatest.funsuite.AnyFunSuite + +class LazyStatsReceiverTest extends AnyFunSuite { + test("Doesn't eagerly initialize counters") { + val underlying = new InMemoryStatsReceiver() + val wrapped = new LazyStatsReceiver(underlying) + val counter = wrapped.counter("foo") + assert(underlying.counters.isEmpty) + counter.incr() + assert(underlying.counters == Map(Seq("foo") -> 1)) + } + + test("Doesn't eagerly initialize histograms") { + val underlying = new InMemoryStatsReceiver() + val wrapped = new LazyStatsReceiver(underlying) + val histo = wrapped.stat("foo") + assert(underlying.stats.isEmpty) + histo.add(5) + assert(underlying.stats == Map(Seq("foo") -> Seq(5))) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/MetricBuilderTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/MetricBuilderTest.scala new file mode 100644 index 0000000000..101a6c327f --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/MetricBuilderTest.scala @@ -0,0 +1,17 @@ +package com.twitter.finagle.stats + +import com.twitter.finagle.stats.MetricBuilder.CounterType +import org.scalatest.funsuite.AnyFunSuite + +class MetricBuilderTest extends AnyFunSuite { + + test("metadataScopeSeparator affects stringification") { + val mb = MetricBuilder(name = Seq("foo", "bar"), metricType = CounterType) + + assert(mb.toString.contains("foo/bar")) + metadataScopeSeparator.setSeparator("-") + assert(mb.toString.contains("foo-bar")) + metadataScopeSeparator.setSeparator("_aTerriblyLongStringForSomeReason_") + assert(mb.toString.contains("foo_aTerriblyLongStringForSomeReason_bar")) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/MultiCategorizingExceptionStatsHandlerTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/MultiCategorizingExceptionStatsHandlerTest.scala index 8bc206cdcb..55ec7d313c 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/MultiCategorizingExceptionStatsHandlerTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/MultiCategorizingExceptionStatsHandlerTest.scala @@ -1,11 +1,9 @@ package com.twitter.finagle.stats -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite +import scala.collection.compat._ -@RunWith(classOf[JUnitRunner]) -class MultiCategorizingExceptionStatsHandlerTest extends FunSuite { +class MultiCategorizingExceptionStatsHandlerTest extends AnyFunSuite { test("uses label, flags, source, exception chain and rolls up") { val receiver = new InMemoryStatsReceiver @@ -22,6 +20,16 @@ class MultiCategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted + val identity = receiver.schemas(Seq("clienterrors")).toMetricBuilder.get.identity + assert( + identity.labels == Map( + "interrupted" -> "true", + "restartable" -> "true", + "source" -> "service", + "exception" -> "java.lang.Exception" + ) + ) + assert(receiver.counters(Seq("clienterrors")) == 1) assert(receiver.counters(Seq("clienterrors", "interrupted")) == 1) assert(receiver.counters(Seq("clienterrors", "restartable")) == 1) @@ -101,6 +109,9 @@ class MultiCategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted + val identity = receiver.schemas(Seq("clienterrors")).toMetricBuilder.get.identity + assert(identity.labels == Map("source" -> "service", "exception" -> "java.lang.Exception")) + assert(receiver.counters(Seq("clienterrors")) == 1) assert(receiver.counters(Seq("clienterrors", classOf[RuntimeException].getName)) == 1) assert( @@ -144,8 +155,11 @@ class MultiCategorizingExceptionStatsHandlerTest extends FunSuite { val keys = receiver.counters.keys.map(_.mkString("/")).toSeq.sorted - assert(receiver.counters.filterKeys(_.contains("failures")).size == 3) - assert(receiver.counters.filterKeys(_.contains("sourcedfailures")).size == 0) + val identity = receiver.schemas(Seq("failures")).toMetricBuilder.get.identity + assert(identity.labels == Map("exception" -> "java.lang.Exception")) + + assert(receiver.counters.view.filterKeys(_.contains("failures")).size == 3) + assert(receiver.counters.view.filterKeys(_.contains("sourcedfailures")).size == 0) assert( keys == Seq( diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiverTest.scala new file mode 100644 index 0000000000..3a3d5eae1c --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/RoleConfiguredStatsReceiverTest.scala @@ -0,0 +1,38 @@ +package com.twitter.finagle.stats + +import org.scalatest.funsuite.AnyFunSuite + +class RoleConfiguredStatsReceiverTest extends AnyFunSuite { + test("RoleConfiguredSR configures metadata role correctly") { + val mem = new InMemoryStatsReceiver + val clientSR = new RoleConfiguredStatsReceiver(mem, SourceRole.Client) + val serverSR = new RoleConfiguredStatsReceiver(mem, SourceRole.Server) + + clientSR.counter("foo") + assert(mem.schemas.get(Seq("foo")).get.role == SourceRole.Client) + clientSR.addGauge("bar")(3) + assert(mem.schemas.get(Seq("bar")).get.role == SourceRole.Client) + clientSR.stat("baz") + assert(mem.schemas.get(Seq("baz")).get.role == SourceRole.Client) + + serverSR.counter("fee") + assert(mem.schemas.get(Seq("fee")).get.role == SourceRole.Server) + serverSR.addGauge("fi")(3) + assert(mem.schemas.get(Seq("fi")).get.role == SourceRole.Server) + serverSR.stat("foe") + assert(mem.schemas.get(Seq("foe")).get.role == SourceRole.Server) + } + + test("RoleConfiguredSR equality") { + val mem = new InMemoryStatsReceiver + val clientA = RoleConfiguredStatsReceiver(mem, SourceRole.Client) + val clientB = RoleConfiguredStatsReceiver(mem, SourceRole.Client) + val serverA = RoleConfiguredStatsReceiver(mem, SourceRole.Server) + val serverB = RoleConfiguredStatsReceiver(mem, SourceRole.Server) + + assert(clientA == clientB) + assert(clientA != serverA) + assert(clientA != serverB) + assert(serverA == serverB) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/StatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/StatsReceiverTest.scala index 76d068d80f..d279075565 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/StatsReceiverTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/StatsReceiverTest.scala @@ -1,17 +1,22 @@ package com.twitter.finagle.stats -import com.twitter.conversions.time._ -import com.twitter.util.{Await, Future} +import com.twitter.conversions.DurationOps._ +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.CounterishGaugeType +import com.twitter.finagle.stats.MetricBuilder.GaugeType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats.MetricBuilder.IdentityType +import com.twitter.finagle.stats.MetricBuilder.UnlatchedCounter +import com.twitter.util.Await +import com.twitter.util.Future import java.util.concurrent.TimeUnit -import org.junit.runner.RunWith import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite import scala.collection.mutable.ArrayBuffer -@RunWith(classOf[JUnitRunner]) -class StatsReceiverTest extends FunSuite { - test("RollupStatsReceiver counter/stats") { +class StatsReceiverTest extends AnyFunSuite { + + test("RollupStatsReceiver counter/stats - full translation") { val mem = new InMemoryStatsReceiver val receiver = new RollupStatsReceiver(mem) @@ -25,12 +30,43 @@ class StatsReceiverTest extends FunSuite { assert(mem.counters(Seq("toto", "titi")) == 2) assert(mem.counters(Seq("toto", "titi", "tata")) == 1) assert(mem.counters(Seq("toto", "titi", "tutu")) == 1) + + def identity(path: String*): IdentityType = mem.schemas(path).identity.identityType + + assert(identity("toto") == IdentityType.NonDeterminate) + assert(identity("toto", "titi") == IdentityType.NonDeterminate) + assert(identity("toto", "titi", "tata") == IdentityType.NonDeterminate) + assert(identity("toto", "titi", "tutu") == IdentityType.NonDeterminate) + } + + test("RollupStatsReceiver counter/stats - hierarchical only") { + val mem = new InMemoryStatsReceiver + val receiver = new RollupStatsReceiver(mem, hierarchicalOnly = true) + + receiver.counter("toto", "titi", "tata").incr() + assert(mem.counters(Seq("toto")) == 1) + assert(mem.counters(Seq("toto", "titi")) == 1) + assert(mem.counters(Seq("toto", "titi", "tata")) == 1) + + receiver.counter("toto", "titi", "tutu").incr() + assert(mem.counters(Seq("toto")) == 2) + assert(mem.counters(Seq("toto", "titi")) == 2) + assert(mem.counters(Seq("toto", "titi", "tata")) == 1) + assert(mem.counters(Seq("toto", "titi", "tutu")) == 1) + + def identity(path: String*): IdentityType = mem.schemas(path).identity.identityType + + assert(identity("toto") == IdentityType.NonDeterminate) + assert(identity("toto", "titi") == IdentityType.HierarchicalOnly) + assert(identity("toto", "titi", "tata") == IdentityType.HierarchicalOnly) + assert(identity("toto", "titi", "tutu") == IdentityType.HierarchicalOnly) } test("Broadcast Counter/Stat") { class MemCounter extends Counter { var c: Long = 0 def incr(delta: Long): Unit = { c += delta } + def metadata: Metadata = NoMetadata } val c1 = new MemCounter val c2 = new MemCounter @@ -43,8 +79,9 @@ class StatsReceiverTest extends FunSuite { assert(c2.c == 1) class MemStat extends Stat { - var values: Seq[Float] = ArrayBuffer.empty[Float] + var values: scala.collection.Seq[Float] = ArrayBuffer.empty[Float] def add(f: Float): Unit = { values = values :+ f } + def metadata: Metadata = NoMetadata } val s1 = new MemStat val s2 = new MemStat @@ -52,9 +89,9 @@ class StatsReceiverTest extends FunSuite { assert(s1.values == Seq.empty) assert(s2.values == Seq.empty) - broadcastStat.add(1F) - assert(s1.values == Seq(1F)) - assert(s2.values == Seq(1F)) + broadcastStat.add(1f) + assert(s1.values == Seq(1f)) + assert(s2.values == Seq(1f)) } test("StatsReceiver time") { @@ -92,6 +129,65 @@ class StatsReceiverTest extends FunSuite { verify(receiver, times(3)).stat("2", "chainz") } + test("StatsReceiver.label: a non-empty label name is required") { + val sr = new InMemoryStatsReceiver + intercept[IllegalArgumentException] { + sr.label("", "cool") + } + } + + test("StatsReceiver.label: empty label values are effectively no-ops") { + val sr = new InMemoryStatsReceiver + sr.label("foo", "").counter("bar").incr() + assert(sr.counters(Seq("bar")) == 1) + assert(sr.schemas(Seq("bar")).identity.labels.isEmpty) + } + + test("StatsReceiver.scope: by default has the same behavior has hierarchicalScope") { + val sr = new InMemoryStatsReceiver + val scopedCounter = sr.scope("cool").counter("numbers") + + val mb = scopedCounter.metadata.toMetricBuilder.get + + assert(mb.identity.labels.isEmpty) + assert(mb.identity.hierarchicalName == Seq("cool", "numbers")) + assert(mb.identity.dimensionalName == Seq("numbers")) + assert(mb.identity.identityType == IdentityType.NonDeterminate) + } + + test("StatsReceiver.hierarchicalScope: composes with label correctly") { + val sr = new InMemoryStatsReceiver + val scopedCounter = sr + .hierarchicalScope("clnt") + .hierarchicalScope("backend") + .label("clnt", "backend") + .counter("requests") + + val mb = scopedCounter.metadata.toMetricBuilder.get + + assert(mb.identity.hierarchicalName == Seq("clnt", "backend", "requests")) + assert(mb.identity.dimensionalName == Seq("requests")) + assert(mb.identity.labels == Map("clnt" -> "backend")) + assert(mb.identity.identityType == IdentityType.NonDeterminate) + } + + test("StatsReceiver.dimensionalScope: composes with label correctly") { + val sr = new InMemoryStatsReceiver + val scopedCounter = sr + .dimensionalScope("foo") + .dimensionalScope("baz") + .hierarchicalScope("bar") + .label("clnt", "backend") + .counter("requests") + + val mb = scopedCounter.metadata.toMetricBuilder.get + + assert(mb.identity.hierarchicalName == Seq("bar", "requests")) + assert(mb.identity.dimensionalName == Seq("foo", "baz", "requests")) + assert(mb.identity.labels == Map("clnt" -> "backend")) + assert(mb.identity.identityType == IdentityType.NonDeterminate) + } + test("StatsReceiver.scope: prefix stats by a scope string") { val receiver = new InMemoryStatsReceiver val scoped = receiver.scope("foo") @@ -143,16 +239,21 @@ class StatsReceiverTest extends FunSuite { assert("NullStatsReceiver" == NullStatsReceiver.scope("hi").scopeSuffix("bye").toString) assert( - "BlacklistStatsReceiver(NullStatsReceiver)" == - new BlacklistStatsReceiver(NullStatsReceiver, { _ => - false - }).toString + "DenylistStatsReceiver(NullStatsReceiver)" == + new DenylistStatsReceiver(NullStatsReceiver, { _ => false }).toString ) val inMem = new InMemoryStatsReceiver() assert("InMemoryStatsReceiver" == inMem.toString) assert("InMemoryStatsReceiver/scope1" == inMem.scope("scope1").toString) + + assert("InMemoryStatsReceiver/scope1" == inMem.hierarchicalScope("scope1").toString) + + assert("InMemoryStatsReceiver" == inMem.dimensionalScope("scope1").toString) + + assert("InMemoryStatsReceiver" == inMem.label("a", "b").toString) + assert( "InMemoryStatsReceiver/scope1/scope2" == inMem.scope("scope1").scope("scope2").toString @@ -180,4 +281,24 @@ class StatsReceiverTest extends FunSuite { } + test("StatsReceiver validate and record metrics") { + val sr = new InMemoryStatsReceiver() + val counter = MetricBuilder(name = Seq("a"), metricType = CounterType) + val counterishGauge = + MetricBuilder(name = Seq("b"), metricType = CounterishGaugeType) + val gauge = MetricBuilder(name = Seq("c"), metricType = GaugeType) + val stat = MetricBuilder(name = Seq("d"), metricType = HistogramType) + val unlatchedCounter = + MetricBuilder(name = Seq("e"), metricType = UnlatchedCounter) + + sr.addGauge(gauge)(1) + sr.addGauge(counterishGauge)(1) + sr.counter(counter) + sr.counter(unlatchedCounter) + sr.stat(stat) + + intercept[IllegalArgumentException] { + sr.counter(counterishGauge) + } + } } diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiverTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiverTest.scala index b85b7203a3..8d73923946 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiverTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/VerbosityAdjustingStatsReceiverTest.scala @@ -1,8 +1,9 @@ package com.twitter.finagle.stats -import org.scalatest.{FunSuite, OneInstancePerTest} +import org.scalatest.OneInstancePerTest +import org.scalatest.funsuite.AnyFunSuite -class VerbosityAdjustingStatsReceiverTest extends FunSuite with OneInstancePerTest { +class VerbosityAdjustingStatsReceiverTest extends AnyFunSuite with OneInstancePerTest { val inMemory = new InMemoryStatsReceiver() val verbose = new VerbosityAdjustingStatsReceiver(inMemory, Verbosity.Debug) diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/WithHistogramDetailsTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/WithHistogramDetailsTest.scala index c730ba8fb0..89261f65e6 100644 --- a/util-stats/src/test/scala/com/twitter/finagle/stats/WithHistogramDetailsTest.scala +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/WithHistogramDetailsTest.scala @@ -1,11 +1,8 @@ package com.twitter.finagle.stats -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class WithHistogramDetailsTest extends FunSuite { +class WithHistogramDetailsTest extends AnyFunSuite { test("AggregateWithHistogramDetails rejects empties") { intercept[IllegalArgumentException] { AggregateWithHistogramDetails(Nil) diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/exp/DefaultExpressionTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/exp/DefaultExpressionTest.scala new file mode 100644 index 0000000000..a7142ec4e9 --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/exp/DefaultExpressionTest.scala @@ -0,0 +1,39 @@ +package com.twitter.finagle.stats.exp + +import com.twitter.finagle.stats.MetricUnit +import com.twitter.finagle.stats.Milliseconds +import com.twitter.finagle.stats.Percentage +import com.twitter.finagle.stats.Requests +import org.scalatest.funsuite.AnyFunSuite + +class DefaultExpressionTest extends AnyFunSuite { + test("success rate expression") { + val successRate = DefaultExpression.successRate(Expression(97), Expression(3)) + assertExpression(successRate, ExpressionNames.successRateName, Percentage) + assert(successRate.bounds == MonotoneThresholds(GreaterThan, 99.5, 99.97)) + } + + test("throughput") { + val throughput = DefaultExpression.throughput(Expression(7777)) + assertExpression(throughput, ExpressionNames.throughputName, Requests) + } + + test("failures") { + val failures = DefaultExpression.failures(Expression(0)) + assertExpression(failures, ExpressionNames.failuresName, Requests) + } + + test("latency p99") { + val latencyP99 = DefaultExpression.latency99(Expression(1000)) + assertExpression(latencyP99, ExpressionNames.latencyName, Milliseconds) + } + + private[this] def assertExpression( + expression: ExpressionSchema, + name: String, + unit: MetricUnit + ): Unit = { + assert(expression.name == name) + assert(expression.unit == unit) + } +} diff --git a/util-stats/src/test/scala/com/twitter/finagle/stats/exp/ExpressionSchemaTest.scala b/util-stats/src/test/scala/com/twitter/finagle/stats/exp/ExpressionSchemaTest.scala new file mode 100644 index 0000000000..15e512aeff --- /dev/null +++ b/util-stats/src/test/scala/com/twitter/finagle/stats/exp/ExpressionSchemaTest.scala @@ -0,0 +1,245 @@ +package com.twitter.finagle.stats.exp + +import com.twitter.conversions.DurationOps._ +import com.twitter.finagle.stats.MetricBuilder.CounterType +import com.twitter.finagle.stats.MetricBuilder.HistogramType +import com.twitter.finagle.stats._ +import com.twitter.finagle.stats.exp.HistogramComponent.Percentile +import com.twitter.util.Stopwatch +import com.twitter.util.Time +import com.twitter.util.TimeControl +import scala.util.control.NonFatal +import org.scalatest.funsuite.AnyFunSuite + +class ExpressionSchemaTest extends AnyFunSuite { + trait Ctx { + val sr = new InMemoryStatsReceiver + val clientSR = RoleConfiguredStatsReceiver(sr, SourceRole.Client, Some("downstream")) + + val successMb = + MetricBuilder(name = Seq("success"), metricType = CounterType) + val failuresMb = + MetricBuilder(name = Seq("failures"), metricType = CounterType) + val latencyMb = + MetricBuilder(name = Seq("latency"), metricType = HistogramType) + val sum = Expression(successMb).plus(Expression(failuresMb)) + + def slowQuery(ctl: TimeControl, succeed: Boolean): Unit = { + ctl.advance(50.milliseconds) + if (!succeed) { + throw new Exception("boom!") + } + } + + val downstreamLabel = + Map( + ExpressionSchema.Role -> SourceRole.Client.toString, + ExpressionSchema.ServiceName -> "downstream") + + protected[this] def nameToKey( + name: String, + labels: Map[String, String] = Map(), + namespaces: Seq[String] = Seq() + ): ExpressionSchemaKey = + ExpressionSchemaKey(name, downstreamLabel ++ labels, namespaces) + } + + test("Demonstrate an end to end example of using metric expressions") { + new Ctx { + val successCounter = clientSR.counter(successMb) + val failuresCounter = clientSR.counter(failuresMb) + val latencyStat = clientSR.stat(latencyMb) + + val successRate = ExpressionSchema("success_rate", Expression(successMb).divide(sum)) + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.75)) + .withUnit(Percentage) + .withDescription("The success rate of the slow query") + + val latency = + ExpressionSchema("latency", Expression(latencyMb, HistogramComponent.Percentile(0.99))) + .withUnit(Milliseconds) + .withDescription("The latency of the slow query") + + clientSR.registerExpression(successRate) + clientSR.registerExpression(latency) + + def runTheQuery(succeed: Boolean): Unit = { + Time.withCurrentTimeFrozen { ctl => + val elapsed = Stopwatch.start() + try { + slowQuery(ctl, succeed) + successCounter.incr() + } catch { + case NonFatal(exn) => + failuresCounter.incr() + } finally { + latencyStat.add(elapsed().inMilliseconds.toFloat) + } + } + } + runTheQuery(true) + runTheQuery(false) + + assert(sr.expressions(nameToKey("success_rate")).name == successRate.name) + assert(sr.expressions(nameToKey("success_rate")).expr == successRate.expr) + assert(sr.expressions(nameToKey("latency", Map("bucket" -> "p99"))).name == latency.name) + assert(sr.expressions(nameToKey("latency", Map("bucket" -> "p99"))).expr == latency.expr) + + assert(sr.counters(Seq("success")) == 1) + assert(sr.counters(Seq("failures")) == 1) + assert(sr.stats(Seq("latency")) == Seq(50, 50)) + + assert( + sr.expressions(nameToKey("success_rate")).labels( + ExpressionSchema.Role) == SourceRole.Client.toString) + } + } + + test("create a histogram expression requires a component") { + new Ctx { + val e = intercept[IllegalArgumentException] { + ExpressionSchema("latency", Expression(latencyMb)) + } + + assert(e.getMessage.contains("provide a component for histogram")) + } + } + + test("histogram expressions come with default labels") { + new Ctx { + clientSR.registerExpression( + ExpressionSchema("latency", Expression(latencyMb, HistogramComponent.Percentile(0.9)))) + clientSR.registerExpression( + ExpressionSchema("latency", Expression(latencyMb, HistogramComponent.Percentile(0.99)))) + clientSR.registerExpression( + ExpressionSchema("latency", Expression(latencyMb, HistogramComponent.Avg))) + + assert(sr.expressions.contains(nameToKey("latency", Map("bucket" -> "p90")))) + assert(sr.expressions.contains(nameToKey("latency", Map("bucket" -> "p99")))) + assert(sr.expressions.contains(nameToKey("latency", Map("bucket" -> "avg")))) + } + } + + test("expressions with different namespaces are all stored") { + new Ctx { + clientSR.registerExpression( + ExpressionSchema("success_rate", Expression(successMb).divide(sum)) + .withNamespace("section1") + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.75)) + .withUnit(Percentage) + .withDescription("The success rate of the slow query")) + + clientSR.registerExpression( + ExpressionSchema("success_rate", Expression(successMb).divide(sum)) + .withNamespace("section2", "a") + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.75)) + .withUnit(Percentage) + .withDescription("The success rate of the slow query")) + + assert(sr.expressions.values.size == 2) + val namespaces = sr.expressions.values.toSeq.map { exprSchema => + exprSchema.namespace + } + assert(namespaces.contains(Seq("section1"))) + assert(namespaces.contains(Seq("section2", "a"))) + } + } + + test("expressions with the same key - name, labels and namespaces") { + new Ctx { + val successRate1 = clientSR.registerExpression( + ExpressionSchema("success", Expression(successMb)) + .withNamespace("section1") + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.75)) + .withUnit(Percentage) + .withDescription("The success rate of the slow query")) + + val successRate2 = clientSR.registerExpression( + ExpressionSchema("success", Expression(successMb).divide(sum)) + .withNamespace("section1") + .withBounds(MonotoneThresholds(GreaterThan, 99.5, 99.75)) + .withUnit(Percentage) + .withDescription("The success rate of the slow query")) + + // only store the first + // the second attempt will be a Throw + assert(sr.expressions.values.size == 1) + assert( + sr.expressions(nameToKey("success", Map(), Seq("section1"))).expr == Expression(successMb)) + assert(successRate1.isReturn) + assert(successRate2.isThrow) + } + } + + test("StringExpression nested in FunctionExpression") { + val loadedSr = LoadedStatsReceiver.self + val sr = new InMemoryStatsReceiver + LoadedStatsReceiver.self = sr + sr.registerExpression( + ExpressionSchema( + "success_rate", + Expression(100).multiply( + Expression(Seq("clnt", "aclient", "success"), isCounter = true).divide( + Expression(Seq("clnt", "aclient", "success"), isCounter = true).plus( + Expression(Seq("clnt", "aclient", "failures"), isCounter = true)))) + ).withRole(SourceRole.Client).withServiceName("aclient")) + + assert(sr.expressions.values.size == 1) + assert( + sr.expressions.contains( + ExpressionSchemaKey( + "success_rate", + Map("role" -> "Client", "service_name" -> "aclient"), + Seq()))) + LoadedStatsReceiver.self = loadedSr + } + + test("use Expression.counter() to create a success rate ExpressionSchema") { + val sr = new InMemoryStatsReceiver + // create a success rate ExpressionSchema + // it is safe to call `.get` since we don't return `None` for now + sr.registerExpression( + ExpressionSchema( + "success_rate", + Expression(100).multiply( + Expression + .counter(sr, Seq("clnt", "aclient", "success")).get.divide(Expression + .counter(sr, Seq("clnt", "aclient", "success")).get.plus( + Expression.counter(sr, Seq("clnt", "aclient", "failures"), showRollup = true).get))) + ).withRole(SourceRole.Client).withServiceName("aclient")) + + assert(sr.expressions.values.size == 1) + assert( + sr.expressions.contains( + ExpressionSchemaKey( + "success_rate", + Map("role" -> "Client", "service_name" -> "aclient"), + Seq()))) + } + + test( + "use Expression.stats() to create a list of ExpressionSchema with default histogram percentiles") { + val sr = new InMemoryStatsReceiver + // create a list histogram ExpressionSchema with default percentiles + val expressions: Seq[Expression] = + Expression.stats(sr, Seq("clnt", "aclient", "latencies")) + expressions.map { expression => + sr.registerExpression( + ExpressionSchema("latencies", expression) + .withRole(SourceRole.Client).withServiceName("aclient")) + } + + assert(sr.expressions.values.size == expressions.size) + for (Percentile(p) <- HistogramComponent.DefaultPercentiles) { + assert( + sr.expressions.contains( + ExpressionSchemaKey( + "latencies", + Map( + "bucket" -> HistogramFormatter.default.labelPercentile(p), + "role" -> "Client", + "service_name" -> "aclient"), + Seq()))) + } + } +} diff --git a/util-test/BUILD b/util-test/BUILD index edd568c335..2cd5045d38 100644 --- a/util-test/BUILD +++ b/util-test/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-test/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-test/src/test/scala", ], diff --git a/util-test/src/main/scala/BUILD b/util-test/src/main/scala/BUILD index c0163b8395..5b48799dc3 100644 --- a/util-test/src/main/scala/BUILD +++ b/util-test/src/main/scala/BUILD @@ -1,17 +1,21 @@ scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, + # This target exists only to act as an aggregator jar. + sources = ["PantsWorkaround.scala"], + platform = "java8", provides = scala_artifact( org = "com.twitter", name = "util-test", repo = artifactory, ), + strict_deps = True, + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-logging/src/main/scala", + "util/util-test/src/main/scala/com/twitter/logging", + "util/util-test/src/main/scala/com/twitter/util/testing", ], exports = [ "3rdparty/jvm/org/scalatest", + "util/util-test/src/main/scala/com/twitter/logging", + "util/util-test/src/main/scala/com/twitter/util/testing", ], ) diff --git a/util-test/src/main/scala/PantsWorkaround.scala b/util-test/src/main/scala/PantsWorkaround.scala new file mode 100644 index 0000000000..9800d15646 --- /dev/null +++ b/util-test/src/main/scala/PantsWorkaround.scala @@ -0,0 +1,4 @@ +// workaround for https://github.com/pantsbuild/pants/issues/6678 +private final class PantsWorkaround { + throw new IllegalStateException("workaround for https://github.com/pantsbuild/pants/issues/6678") +} diff --git a/util-test/src/main/scala/com/twitter/logging/BUILD b/util-test/src/main/scala/com/twitter/logging/BUILD new file mode 100644 index 0000000000..8c33560d6c --- /dev/null +++ b/util-test/src/main/scala/com/twitter/logging/BUILD @@ -0,0 +1,18 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-test-logging", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/scalatest", + "util/util-logging/src/main/scala/com/twitter/logging", + ], + exports = [ + "3rdparty/jvm/org/scalatest", + ], +) diff --git a/util-test/src/main/scala/com/twitter/logging/TestLogging.scala b/util-test/src/main/scala/com/twitter/logging/TestLogging.scala index b877e8e66e..7a741bc1d8 100644 --- a/util-test/src/main/scala/com/twitter/logging/TestLogging.scala +++ b/util-test/src/main/scala/com/twitter/logging/TestLogging.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -18,13 +18,15 @@ package com.twitter.logging import java.util.{logging => jlogging} -import org.scalatest.{BeforeAndAfter, WordSpec} +import org.scalatest.BeforeAndAfter +import org.scalatest.funsuite.AnyFunSuite /** * Specify logging during unit tests via system property, defaulting to FATAL only. */ -trait TestLogging extends BeforeAndAfter { self: WordSpec => - val logLevel = +@deprecated("Prefer using slf4j for logging", "2019-11-20") +trait TestLogging extends BeforeAndAfter { self: AnyFunSuite => + val logLevel: Level = Logger.levelNames(Option[String](System.getenv("log")).getOrElse("FATAL").toUpperCase) private val logger = Logger.get("") @@ -66,13 +68,13 @@ trait TestLogging extends BeforeAndAfter { self: WordSpec => logger.addHandler(traceHandler) } - def logLines(): Seq[String] = traceHandler.get.split("\n") + def logLines(): Seq[String] = traceHandler.get.split("\n").toSeq /** * Verify that the logger set up with `traceLogger` has received a log line with the given * substring somewhere inside it. */ - def mustLog(substring: String) = { + def mustLog(substring: String): org.scalatest.Assertion = { assert(logLines().filter { _ contains substring }.size > 0) } } diff --git a/util-test/src/main/scala/com/twitter/util/testing/ArgumentCapture.scala b/util-test/src/main/scala/com/twitter/util/testing/ArgumentCapture.scala index 21172d68c7..9961cf33ba 100644 --- a/util-test/src/main/scala/com/twitter/util/testing/ArgumentCapture.scala +++ b/util-test/src/main/scala/com/twitter/util/testing/ArgumentCapture.scala @@ -2,7 +2,7 @@ package com.twitter.util.testing import org.mockito.ArgumentCaptor import org.mockito.exceptions.Reporter -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ import scala.reflect._ // This file was generated from codegen/util-test/ArgumentCapture.scala.mako @@ -257,7 +257,9 @@ trait ArgumentCapture { E: ClassTag, F: ClassTag, G: ClassTag - ](func: (A, B, C, D, E, F, G) => _): Seq[(A, B, C, D, E, F, G)] = { + ]( + func: (A, B, C, D, E, F, G) => _ + ): Seq[(A, B, C, D, E, F, G)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -293,7 +295,9 @@ trait ArgumentCapture { E: ClassTag, F: ClassTag, G: ClassTag - ](func: (A, B, C, D, E, F, G) => _): (A, B, C, D, E, F, G) = + ]( + func: (A, B, C, D, E, F, G) => _ + ): (A, B, C, D, E, F, G) = capturingAll[A, B, C, D, E, F, G](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 8 iterables together into a Seq of 8-tuples. */ @@ -329,7 +333,9 @@ trait ArgumentCapture { F: ClassTag, G: ClassTag, H: ClassTag - ](func: (A, B, C, D, E, F, G, H) => _): Seq[(A, B, C, D, E, F, G, H)] = { + ]( + func: (A, B, C, D, E, F, G, H) => _ + ): Seq[(A, B, C, D, E, F, G, H)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -369,7 +375,9 @@ trait ArgumentCapture { F: ClassTag, G: ClassTag, H: ClassTag - ](func: (A, B, C, D, E, F, G, H) => _): (A, B, C, D, E, F, G, H) = + ]( + func: (A, B, C, D, E, F, G, H) => _ + ): (A, B, C, D, E, F, G, H) = capturingAll[A, B, C, D, E, F, G, H](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 9 iterables together into a Seq of 9-tuples. */ @@ -408,7 +416,9 @@ trait ArgumentCapture { G: ClassTag, H: ClassTag, I: ClassTag - ](func: (A, B, C, D, E, F, G, H, I) => _): Seq[(A, B, C, D, E, F, G, H, I)] = { + ]( + func: (A, B, C, D, E, F, G, H, I) => _ + ): Seq[(A, B, C, D, E, F, G, H, I)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -452,7 +462,9 @@ trait ArgumentCapture { G: ClassTag, H: ClassTag, I: ClassTag - ](func: (A, B, C, D, E, F, G, H, I) => _): (A, B, C, D, E, F, G, H, I) = + ]( + func: (A, B, C, D, E, F, G, H, I) => _ + ): (A, B, C, D, E, F, G, H, I) = capturingAll[A, B, C, D, E, F, G, H, I](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 10 iterables together into a Seq of 10-tuples. */ @@ -478,8 +490,9 @@ trait ArgumentCapture { .zip(arg7) .zip(arg8) .zip(arg9) - .map({ case (((((((((a, b), c), d), e), f), g), h), i), j) => (a, b, c, d, e, f, g, h, i, j) } - ) + .map({ + case (((((((((a, b), c), d), e), f), g), h), i), j) => (a, b, c, d, e, f, g, h, i, j) + }) .toSeq } @@ -495,7 +508,9 @@ trait ArgumentCapture { H: ClassTag, I: ClassTag, J: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J) => _): Seq[(A, B, C, D, E, F, G, H, I, J)] = { + ]( + func: (A, B, C, D, E, F, G, H, I, J) => _ + ): Seq[(A, B, C, D, E, F, G, H, I, J)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -543,7 +558,9 @@ trait ArgumentCapture { H: ClassTag, I: ClassTag, J: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J) => _): (A, B, C, D, E, F, G, H, I, J) = + ]( + func: (A, B, C, D, E, F, G, H, I, J) => _ + ): (A, B, C, D, E, F, G, H, I, J) = capturingAll[A, B, C, D, E, F, G, H, I, J](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 11 iterables together into a Seq of 11-tuples. */ @@ -591,7 +608,9 @@ trait ArgumentCapture { I: ClassTag, J: ClassTag, K: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J, K) => _): Seq[(A, B, C, D, E, F, G, H, I, J, K)] = { + ]( + func: (A, B, C, D, E, F, G, H, I, J, K) => _ + ): Seq[(A, B, C, D, E, F, G, H, I, J, K)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -643,7 +662,9 @@ trait ArgumentCapture { I: ClassTag, J: ClassTag, K: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J, K) => _): (A, B, C, D, E, F, G, H, I, J, K) = + ]( + func: (A, B, C, D, E, F, G, H, I, J, K) => _ + ): (A, B, C, D, E, F, G, H, I, J, K) = capturingAll[A, B, C, D, E, F, G, H, I, J, K](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 12 iterables together into a Seq of 12-tuples. */ @@ -694,7 +715,9 @@ trait ArgumentCapture { J: ClassTag, K: ClassTag, L: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J, K, L) => _): Seq[(A, B, C, D, E, F, G, H, I, J, K, L)] = { + ]( + func: (A, B, C, D, E, F, G, H, I, J, K, L) => _ + ): Seq[(A, B, C, D, E, F, G, H, I, J, K, L)] = { val argCaptorA = ArgumentCaptor.forClass(classTag[A].runtimeClass.asInstanceOf[Class[A]]) val argCaptorB = ArgumentCaptor.forClass(classTag[B].runtimeClass.asInstanceOf[Class[B]]) val argCaptorC = ArgumentCaptor.forClass(classTag[C].runtimeClass.asInstanceOf[Class[C]]) @@ -750,7 +773,9 @@ trait ArgumentCapture { J: ClassTag, K: ClassTag, L: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J, K, L) => _): (A, B, C, D, E, F, G, H, I, J, K, L) = + ]( + func: (A, B, C, D, E, F, G, H, I, J, K, L) => _ + ): (A, B, C, D, E, F, G, H, I, J, K, L) = capturingAll[A, B, C, D, E, F, G, H, I, J, K, L](func).lastOption.getOrElse(noArgWasCaptured()) /** Zip 13 iterables together into a Seq of 13-tuples. */ @@ -866,7 +891,9 @@ trait ArgumentCapture { K: ClassTag, L: ClassTag, M: ClassTag - ](func: (A, B, C, D, E, F, G, H, I, J, K, L, M) => _): (A, B, C, D, E, F, G, H, I, J, K, L, M) = + ]( + func: (A, B, C, D, E, F, G, H, I, J, K, L, M) => _ + ): (A, B, C, D, E, F, G, H, I, J, K, L, M) = capturingAll[A, B, C, D, E, F, G, H, I, J, K, L, M](func).lastOption .getOrElse(noArgWasCaptured()) @@ -1523,8 +1550,8 @@ trait ArgumentCapture { .zip(arg17) .map({ case ( - ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), - r + ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), + r ) => (a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r) }) @@ -1701,8 +1728,10 @@ trait ArgumentCapture { .zip(arg18) .map({ case ( - (((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), r), - s + ( + ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), + r), + s ) => (a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s) }) @@ -1887,14 +1916,14 @@ trait ArgumentCapture { .zip(arg19) .map({ case ( - ( ( - ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), - r + ( + ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), + r + ), + s ), - s - ), - t + t ) => (a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s, t) }) @@ -2087,17 +2116,19 @@ trait ArgumentCapture { .zip(arg20) .map({ case ( - ( ( ( - ((((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), q), - r + ( + ( + (((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), + q), + r + ), + s ), - s + t ), - t - ), - u + u ) => (a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s, t, u) }) @@ -2298,23 +2329,25 @@ trait ArgumentCapture { .zip(arg21) .map({ case ( - ( ( ( ( ( - (((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), p), - q + ( + ( + ((((((((((((((a, b), c), d), e), f), g), h), i), j), k), l), m), n), o), + p), + q + ), + r ), - r + s ), - s + t ), - t + u ), - u - ), - v + v ) => (a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s, t, u, v) }) diff --git a/util-test/src/main/scala/com/twitter/util/testing/BUILD b/util-test/src/main/scala/com/twitter/util/testing/BUILD new file mode 100644 index 0000000000..efd4cb9917 --- /dev/null +++ b/util-test/src/main/scala/com/twitter/util/testing/BUILD @@ -0,0 +1,42 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-test-testing", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/scalatest", + "util/util-stats/src/main/scala", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", + ], + exports = [ + "3rdparty/jvm/org/scalatest", + ], +) + +scala_library( + name = "no-mockito", + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-test-testing-no-mockito", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + sources = ["*.scala", "!ArgumentCapture.scala"], + dependencies = [ + "3rdparty/jvm/org/scalatest", + "util/util-stats/src/main/scala", + "util/util-stats/src/main/scala/com/twitter/finagle/stats", + ], + exports = [ + "3rdparty/jvm/org/scalatest", + ], +) diff --git a/util-test/src/main/scala/com/twitter/util/testing/ExpressionTestMixin.scala b/util-test/src/main/scala/com/twitter/util/testing/ExpressionTestMixin.scala new file mode 100644 index 0000000000..2d487a6c43 --- /dev/null +++ b/util-test/src/main/scala/com/twitter/util/testing/ExpressionTestMixin.scala @@ -0,0 +1,60 @@ +package com.twitter.util.testing + +import com.twitter.finagle.stats.exp.ExpressionSchema +import com.twitter.finagle.stats.exp.ExpressionSchemaKey +import com.twitter.finagle.stats.InMemoryStatsReceiver +import org.scalatest.BeforeAndAfterEach +import org.scalatest.Suite +import org.scalatest.SuiteMixin +import scala.collection.mutable + +/** + * Testing trait which provides utilities for testing for proper + * [[com.twitter.finagle.stats.exp.Expression]] registration + */ +trait ExpressionTestMixin extends SuiteMixin with BeforeAndAfterEach { this: Suite => + + protected def downstreamLabel: Map[String, String] + protected val sr: InMemoryStatsReceiver + + protected def nameToKey( + name: String, + labels: Map[String, String] = Map(), + namespaces: Seq[String] = Seq() + ): ExpressionSchemaKey = + ExpressionSchemaKey(name, downstreamLabel ++ labels, namespaces) + + override protected def beforeEach(): Unit = { + sr.clear() + super.beforeEach() + } + + /** + * Given the name, labels, and namespaces that make up a given ExpressionSchemaKey, + * assert that expression is registered on the StatsReceiver. + */ + protected def assertExpressionIsRegistered( + name: String, + labels: Map[String, String] = Map(), + namespaces: Seq[String] = Seq() + ): Unit = { + assert(sr.expressions.contains(nameToKey(name, labels, namespaces))) + } + + /** + * Returns a collection of all registered expressions on the provided StatsReceiver + */ + protected def getMetricsReview: mutable.Map[ExpressionSchemaKey, ExpressionSchema] = { + sr.expressions + } + + /** + * Given a collection of expected ExpressionSchemaKeys, ensure that collection matches what is + * registered. + * + * @param expected a Set of [[ExpressionSchemaKey]]s that should be registered on the StatsReceiver + */ + protected def assertExpressionsAsExpected(expected: Set[ExpressionSchemaKey]): Unit = { + assert(sr.expressions.keySet == expected) + } +} diff --git a/util-test/src/test/scala/BUILD b/util-test/src/test/scala/BUILD index 7d3ae63b7f..4869e39f9f 100644 --- a/util-test/src/test/scala/BUILD +++ b/util-test/src/test/scala/BUILD @@ -1,10 +1,6 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-test/src/main/scala", + "util/util-test/src/test/scala/com/twitter/util/testing", ], ) diff --git a/util-test/src/test/scala/com/twitter/util/testing/ArgumentCaptureTest.scala b/util-test/src/test/scala/com/twitter/util/testing/ArgumentCaptureTest.scala index 1b43cb9c8a..502350d755 100644 --- a/util-test/src/test/scala/com/twitter/util/testing/ArgumentCaptureTest.scala +++ b/util-test/src/test/scala/com/twitter/util/testing/ArgumentCaptureTest.scala @@ -1,14 +1,11 @@ package com.twitter.util.testing -import org.junit.runner.RunWith import org.mockito.Matchers._ import org.mockito.Mockito._ -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ArgumentCaptureTest extends FunSuite with MockitoSugar with ArgumentCapture { +class ArgumentCaptureTest extends AnyFunSuite with MockitoSugar with ArgumentCapture { class MockSubject { def method1(arg: String): Long = 0L def method2(arg1: String, arg2: String): Long = 0L diff --git a/util-test/src/test/scala/com/twitter/util/testing/BUILD b/util-test/src/test/scala/com/twitter/util/testing/BUILD new file mode 100644 index 0000000000..5d09aa73e9 --- /dev/null +++ b/util-test/src/test/scala/com/twitter/util/testing/BUILD @@ -0,0 +1,18 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/mockito:mockito-all", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito-1-10", + "util/util-test/src/main/scala/com/twitter/util/testing", + ], +) diff --git a/util-thrift/BUILD b/util-thrift/BUILD index ce55cabe80..763f4e28d7 100644 --- a/util-thrift/BUILD +++ b/util-thrift/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-thrift/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-thrift/src/test/scala", ], diff --git a/util-thrift/src/main/scala/BUILD b/util-thrift/src/main/scala/BUILD index 01e5d0a011..530a127a84 100644 --- a/util-thrift/src/main/scala/BUILD +++ b/util-thrift/src/main/scala/BUILD @@ -1,19 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-thrift", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", - "3rdparty/jvm/org/apache/thrift:libthrift", - "util/util-codec", - "util/util-core/src/main/scala", - ], - exports = [ - "3rdparty/jvm/org/apache/thrift:libthrift", - "util/util-codec", + "util/util-thrift/src/main/scala/com/twitter/util", ], ) diff --git a/util-thrift/src/main/scala/com/twitter/util/BUILD b/util-thrift/src/main/scala/com/twitter/util/BUILD new file mode 100644 index 0000000000..4187c8639a --- /dev/null +++ b/util-thrift/src/main/scala/com/twitter/util/BUILD @@ -0,0 +1,22 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-thrift", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/org/apache/thrift:libthrift", + "util/util-codec/src/main/scala/com/twitter/util", + "util/util-core:util-core-util", + ], + exports = [ + "3rdparty/jvm/org/apache/thrift:libthrift", + "util/util-codec/src/main/scala/com/twitter/util", + "util/util-core:util-core-util", + ], +) diff --git a/util-thrift/src/main/scala/com/twitter/util/ThriftCodec.scala b/util-thrift/src/main/scala/com/twitter/util/ThriftCodec.scala index 80e97c7ce0..6fc3d4ad12 100644 --- a/util-thrift/src/main/scala/com/twitter/util/ThriftCodec.scala +++ b/util-thrift/src/main/scala/com/twitter/util/ThriftCodec.scala @@ -2,21 +2,22 @@ package com.twitter.util import org.apache.thrift.TBase import org.apache.thrift.protocol.{TBinaryProtocol, TCompactProtocol, TProtocolFactory} +import scala.reflect.{ClassTag, classTag} object ThriftCodec { - def apply[T <: TBase[_, _]: Manifest, P <: TProtocolFactory: Manifest]: ThriftCodec[T, P] = + def apply[T <: TBase[_, _]: ClassTag, P <: TProtocolFactory: ClassTag]: ThriftCodec[T, P] = new ThriftCodec[T, P] } -class ThriftCodec[T <: TBase[_, _]: Manifest, P <: TProtocolFactory: Manifest] +class ThriftCodec[T <: TBase[_, _]: ClassTag, P <: TProtocolFactory: ClassTag] extends Codec[T, Array[Byte]] with ThriftSerializer { protected lazy val prototype: T = - manifest[T].runtimeClass.asInstanceOf[Class[T]].newInstance + classTag[T].runtimeClass.asInstanceOf[Class[T]].newInstance lazy val protocolFactory: TProtocolFactory = - manifest[P].runtimeClass.asInstanceOf[Class[P]].newInstance + classTag[P].runtimeClass.asInstanceOf[Class[P]].newInstance def encode(item: T): Array[Byte] = toBytes(item) @@ -27,7 +28,7 @@ class ThriftCodec[T <: TBase[_, _]: Manifest, P <: TProtocolFactory: Manifest] } } -class BinaryThriftCodec[T <: TBase[_, _]: Manifest] extends ThriftCodec[T, TBinaryProtocol.Factory] +class BinaryThriftCodec[T <: TBase[_, _]: ClassTag] extends ThriftCodec[T, TBinaryProtocol.Factory] -class CompactThriftCodec[T <: TBase[_, _]: Manifest] +class CompactThriftCodec[T <: TBase[_, _]: ClassTag] extends ThriftCodec[T, TCompactProtocol.Factory] diff --git a/util-thrift/src/main/scala/com/twitter/util/ThriftSerializer.scala b/util-thrift/src/main/scala/com/twitter/util/ThriftSerializer.scala index 4f473a11d4..c327dc93d5 100644 --- a/util-thrift/src/main/scala/com/twitter/util/ThriftSerializer.scala +++ b/util-thrift/src/main/scala/com/twitter/util/ThriftSerializer.scala @@ -33,10 +33,10 @@ trait ThriftSerializer extends StringEncoder { } /** - * A thread-safe [[ThriftSerializer]] that uses [[TSimpleJSONProtocol]]. + * A thread-safe [[ThriftSerializer]] that uses `TSimpleJSONProtocol`. */ class JsonThriftSerializer extends ThriftSerializer { - override def protocolFactory = new TSimpleJSONProtocol.Factory + override def protocolFactory: TProtocolFactory = new TSimpleJSONProtocol.Factory override def fromBytes(obj: TBase[_, _], bytes: Array[Byte]): Unit = { val binarySerializer = new BinaryThriftSerializer @@ -46,18 +46,18 @@ class JsonThriftSerializer extends ThriftSerializer { } /** - * A thread-safe [[ThriftSerializer]] that uses [[TBinaryProtocol]]. + * A thread-safe [[ThriftSerializer]] that uses `TBinaryProtocol`. * * @note an implementation using `com.twitter.finagle.thrift.Protocols.binaryFactory` * instead of this is recommended. */ class BinaryThriftSerializer extends ThriftSerializer with Base64StringEncoder { - override def protocolFactory = new TBinaryProtocol.Factory + override def protocolFactory: TProtocolFactory = new TBinaryProtocol.Factory } /** - * A thread-safe [[ThriftSerializer]] that uses [[TCompactProtocol]]. + * A thread-safe [[ThriftSerializer]] that uses `TCompactProtocol`. */ class CompactThriftSerializer extends ThriftSerializer with Base64StringEncoder { - override def protocolFactory = new TCompactProtocol.Factory + override def protocolFactory: TProtocolFactory = new TCompactProtocol.Factory } diff --git a/util-thrift/src/test/java/BUILD b/util-thrift/src/test/java/BUILD deleted file mode 100644 index 3824246be9..0000000000 --- a/util-thrift/src/test/java/BUILD +++ /dev/null @@ -1,7 +0,0 @@ -junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/org/apache/thrift:libthrift", - ], -) diff --git a/util-thrift/src/test/java/com/twitter/util/BUILD b/util-thrift/src/test/java/com/twitter/util/BUILD new file mode 100644 index 0000000000..5f4374f0de --- /dev/null +++ b/util-thrift/src/test/java/com/twitter/util/BUILD @@ -0,0 +1,9 @@ +java_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/apache/thrift:libthrift", + ], +) diff --git a/util-thrift/src/test/java/com/twitter/util/TestThriftStructure.java b/util-thrift/src/test/java/com/twitter/util/TestThriftStructure.java index 5decf718a7..1ec5e84e87 100644 --- a/util-thrift/src/test/java/com/twitter/util/TestThriftStructure.java +++ b/util-thrift/src/test/java/com/twitter/util/TestThriftStructure.java @@ -159,7 +159,7 @@ public void unsetAString() { this.aString = null; } - /** Returns true if field aString is set (has been asigned a value) and false otherwise */ + /** Returns true if field aString is set (has been assigned a value) and false otherwise */ public boolean isSetAString() { return this.aString != null; } @@ -184,7 +184,7 @@ public void unsetANumber() { __isset_bit_vector.clear(__ANUMBER_ISSET_ID); } - /** Returns true if field aNumber is set (has been asigned a value) and false otherwise */ + /** Returns true if field aNumber is set (has been assigned a value) and false otherwise */ public boolean isSetANumber() { return __isset_bit_vector.get(__ANUMBER_ISSET_ID); } @@ -226,7 +226,7 @@ public Object getFieldValue(_Fields field) { throw new IllegalStateException(); } - /** Returns true if field corresponding to fieldID is set (has been asigned a value) and false otherwise */ + /** Returns true if field corresponding to fieldID is set (has been assigned a value) and false otherwise */ public boolean isSet(_Fields field) { if (field == null) { throw new IllegalArgumentException(); diff --git a/util-thrift/src/test/scala/BUILD b/util-thrift/src/test/scala/BUILD index de9be2e063..f43bc4689c 100644 --- a/util-thrift/src/test/scala/BUILD +++ b/util-thrift/src/test/scala/BUILD @@ -1,11 +1,6 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-thrift/src/main/scala", - "util/util-thrift/src/test/java", + "util/util-thrift/src/test/scala/com/twitter/util", ], ) diff --git a/util-thrift/src/test/scala/com/twitter/util/BUILD b/util-thrift/src/test/scala/com/twitter/util/BUILD new file mode 100644 index 0000000000..4b320f6440 --- /dev/null +++ b/util-thrift/src/test/scala/com/twitter/util/BUILD @@ -0,0 +1,17 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-thrift/src/main/scala/com/twitter/util", + "util/util-thrift/src/test/java/com/twitter/util", + ], +) diff --git a/util-thrift/src/test/scala/com/twitter/util/ThriftCodecTest.scala b/util-thrift/src/test/scala/com/twitter/util/ThriftCodecTest.scala index 4cce8c44b1..6e1534f37f 100644 --- a/util-thrift/src/test/scala/com/twitter/util/ThriftCodecTest.scala +++ b/util-thrift/src/test/scala/com/twitter/util/ThriftCodecTest.scala @@ -1,11 +1,8 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.FunSuite -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ThriftCodecTest extends FunSuite { +class ThriftCodecTest extends AnyFunSuite { private def roundTrip(codec: ThriftCodec[TestThriftStructure, _]): Unit = { val struct = new TestThriftStructure("aString", 5) diff --git a/util-thrift/src/test/scala/com/twitter/util/ThriftSerializerTest.scala b/util-thrift/src/test/scala/com/twitter/util/ThriftSerializerTest.scala index e5d50ec307..ae1fc63840 100644 --- a/util-thrift/src/test/scala/com/twitter/util/ThriftSerializerTest.scala +++ b/util-thrift/src/test/scala/com/twitter/util/ThriftSerializerTest.scala @@ -5,7 +5,7 @@ * 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 + * https://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, @@ -16,12 +16,9 @@ package com.twitter.util -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ThriftSerializerTest extends WordSpec { +class ThriftSerializerTest extends AnyFunSuite { val aString = "me gustan los tacos y los burritos" val aNumber = 42 val original = new TestThriftStructure(aString, aNumber) @@ -48,17 +45,15 @@ class ThriftSerializerTest extends WordSpec { true } - "ThriftSerializer" should { - "encode and decode json" in { - testSerializer(new JsonThriftSerializer, Some(json)) - } + test("ThriftSerializer should encode and decode json") { + testSerializer(new JsonThriftSerializer, Some(json)) + } - "encode and decode binary" in { - testSerializer(new BinaryThriftSerializer, encodedBinary) - } + test("ThriftSerializer should encode and decode binary") { + testSerializer(new BinaryThriftSerializer, encodedBinary) + } - "encode and decode compact" in { - testSerializer(new CompactThriftSerializer, encodedCompact) - } + test("ThriftSerializer should encode and decode compact") { + testSerializer(new CompactThriftSerializer, encodedCompact) } } diff --git a/util-tunable/BUILD b/util-tunable/BUILD index bd6e06dc94..354c2a2f49 100644 --- a/util-tunable/BUILD +++ b/util-tunable/BUILD @@ -1,4 +1,5 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-tunable/src/main/scala", ], diff --git a/util-tunable/lint-configuration b/util-tunable/lint-configuration index 5966217244..8c8021f066 100755 --- a/util-tunable/lint-configuration +++ b/util-tunable/lint-configuration @@ -26,4 +26,4 @@ do done args=$(for x in "$@"; do echo "$x"; done) -./pants run --quiet util/util-tunable/src/main/scala/com/twitter/util/tunable/linter:configuration-linter -- $args \ No newline at end of file +./bazel run --ui_event_filters=-info,-stdout,-stderr --noshow_progress util/util-tunable/src/main/scala/com/twitter/util/tunable/linter:configuration-linter -- $args diff --git a/util-tunable/src/main/scala/BUILD b/util-tunable/src/main/scala/BUILD index b6dde0cb2b..26c9ca7551 100644 --- a/util-tunable/src/main/scala/BUILD +++ b/util-tunable/src/main/scala/BUILD @@ -1,23 +1,6 @@ -scala_library( - name = "scala", - sources = rglobs( - "*.scala", - exclude = [rglobs("util/tunable/linter/*")], - ), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-tunable", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", - "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", - "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - ], - exports = [ - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "util/util-tunable/src/main/scala/com/twitter/util/tunable", ], ) diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/BUILD b/util-tunable/src/main/scala/com/twitter/util/tunable/BUILD new file mode 100644 index 0000000000..aa8915afba --- /dev/null +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/BUILD @@ -0,0 +1,23 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-tunable", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-annotations", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "util/util-app/src/main/scala", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + ], + exports = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + ], +) diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/Deserializers.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/Deserializers.scala index 9d6a31a55a..848d5bceb9 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/Deserializers.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/Deserializers.scala @@ -14,17 +14,19 @@ private[twitter] object json { * Deserialize a String representation of a [[com.twitter.util.Duration]] using * [[com.twitter.util.Duration.parse]] */ - val DurationFromString = new StdDeserializer[util.Duration](classOf[util.Duration]) { - override def deserialize( - jsonParser: JsonParser, - deserializationContext: DeserializationContext - ): util.Duration = util.Duration.parse(jsonParser.getText) - } + val DurationFromString: StdDeserializer[util.Duration] = + new StdDeserializer[util.Duration](classOf[util.Duration]) { + override def deserialize( + jsonParser: JsonParser, + deserializationContext: DeserializationContext + ): util.Duration = util.Duration.parse(jsonParser.getText) + } - val StorageUnitFromString = new StdDeserializer[util.StorageUnit](classOf[util.StorageUnit]) { - override def deserialize( - jsonParser: JsonParser, - deserializationContext: DeserializationContext - ): util.StorageUnit = util.StorageUnit.parse(jsonParser.getText) - } + val StorageUnitFromString: StdDeserializer[util.StorageUnit] = + new StdDeserializer[util.StorageUnit](classOf[util.StorageUnit]) { + override def deserialize( + jsonParser: JsonParser, + deserializationContext: DeserializationContext + ): util.StorageUnit = util.StorageUnit.parse(jsonParser.getText) + } } diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/JsonTunableMapper.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/JsonTunableMapper.scala index 746ccc53b4..ff3a7add75 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/JsonTunableMapper.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/JsonTunableMapper.scala @@ -6,7 +6,7 @@ import com.fasterxml.jackson.databind.{JsonDeserializer, ObjectMapper} import com.fasterxml.jackson.module.scala.DefaultScalaModule import com.twitter.util.{Return, Throw, Try} import java.net.URL -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ object JsonTunableMapper { @@ -15,19 +15,19 @@ object JsonTunableMapper { private case class JsonTunable( @JsonProperty(required = true) id: String, @JsonProperty(value = "type", required = true) valueType: Class[Any], - @JsonProperty(required = true) value: Any - ) + @JsonProperty(required = true) value: Any, + @JsonProperty(required = false) comment: String) private case class JsonTunables(@JsonProperty(required = true) tunables: Seq[JsonTunable]) /** * The Deserializers that [[JsonTunableMapper]] uses by default, in addition to Scala data type - * deserializers afforded by [[com.fasterxml.jackson.module.scala.DefaultScalaModule]]. + * deserializers afforded by `com.fasterxml.jackson.module.scala.DefaultScalaModule`. * * These deserializers are: * - * - [[com.twitter.util.tunable.json.DurationFromString]] - * - [[com.twitter.util.tunable.json.StorageUnitFromString]] + * - `com.twitter.util.tunable.json.DurationFromString` + * - `com.twitter.util.tunable.json.StorageUnitFromString` */ val DefaultDeserializers: Seq[JsonDeserializer[_]] = Seq(DurationFromString, StorageUnitFromString) @@ -49,10 +49,10 @@ object JsonTunableMapper { * in priority order. Where environementOpt and instanceOpt are available (env, instance), * paths are ordered: * - * i. $root/$env/instance-$id.json - * i. $root/$env/instances.json - * i. $root/instance-$id.json - * i. $root/instances.json + * i. \$root/\$env/instance-\$id.json + * i. \$root/\$env/instances.json + * i. \$root/instance-\$id.json + * i. \$root/instances.json */ def pathsByPriority( root: String, @@ -82,36 +82,35 @@ object JsonTunableMapper { /** * Parses a given JSON string into a [[TunableMap]]. The expected format is: * - * { + * {{{ * "tunables": * [ * { - * "id" : "$id1", - * "value" : $value, - * "type" : "$class" + * "id" : "\$id1", + * "value" : \$value, + * "type" : "\$class" * }, * { - * "id" : "$id2", - * "value" : $value, - * "type" : "$class" + * "id" : "\$id2", + * "value" : \$value, + * "type" : "\$class", + * "comment": "optional comment" * } * ] - * } + * }}} * - * Where $id1 and $id2 are unique identifiers used to access the [[Tunable]], $value is the value, - * and $class is the fully-qualified class name (e.g. com.twitter.util.Duration) + * Where \$id1 and \$id2 are unique identifiers used to access the [[Tunable]], \$value is the value, + * and \$class is the fully-qualified class name (e.g. com.twitter.util.Duration) * * If the JSON is invalid, or contains duplicate ids for [[Tunable]]s, `parse` will - * return a [[Throw]]. Otherwise, `parse` returns [[Return[TunableMap]] + * return a [[Throw]]. Otherwise, `parse` returns [[Return Return[TunableMap]] */ final class JsonTunableMapper(deserializers: Seq[JsonDeserializer[_ <: Any]]) { import JsonTunableMapper._ private[this] object DeserializationModule extends SimpleModule { - deserializers.foreach { jd => - addDeserializer(jd.handledType().asInstanceOf[Class[Any]], jd) - } + deserializers.foreach { jd => addDeserializer(jd.handledType().asInstanceOf[Class[Any]], jd) } } private[this] val mapper: ObjectMapper = @@ -172,8 +171,8 @@ final class JsonTunableMapper(deserializers: Seq[JsonDeserializer[_ <: Any]]) { * Load and parse the JSON file located at `path` in the application's resources. * * If no configuration files exist, return [[NullTunableMap]]. - * If multiple configuration files exists, return [[IllegalArgumentException]] - * If the configuration file cannot be parsed, return [[IllegalArgumentException]] + * If multiple configuration files exists, return `IllegalArgumentException` + * If the configuration file cannot be parsed, return `IllegalArgumentException` */ def loadJsonTunables(id: String, path: String): TunableMap = { val classLoader = getClass.getClassLoader diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/ServiceLoadedTunableMap.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/ServiceLoadedTunableMap.scala index 2b7c88e549..81386a3c56 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/ServiceLoadedTunableMap.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/ServiceLoadedTunableMap.scala @@ -1,6 +1,9 @@ package com.twitter.util.tunable import com.twitter.app.LoadService +import com.twitter.concurrent.Once +import java.util.concurrent.ConcurrentHashMap +import java.util.function.BiFunction /** * [[TunableMap]] loaded through `com.twitter.app.LoadService`. @@ -16,6 +19,8 @@ trait ServiceLoadedTunableMap extends TunableMap { object ServiceLoadedTunableMap { + private[this] val loadedMaps = new ConcurrentHashMap[String, TunableMap]() + /** * Uses `com.twitter.app.LoadService` to load the [[TunableMap]] for a given `id`. * @@ -24,18 +29,46 @@ object ServiceLoadedTunableMap { * If more than one match is found, an `IllegalStateException` is thrown. */ def apply(id: String): TunableMap = { + init() + loadedMaps.getOrDefault(id, NullTunableMap) + } - val tunableMaps = LoadService[ServiceLoadedTunableMap]().filter(_.id == id).toList - - tunableMaps match { - case Nil => - NullTunableMap - case tunableMap :: Nil => - tunableMap - case _ => - throw new IllegalStateException( - s"Found multiple `ServiceLoadedTunableMap`s for $id: ${tunableMaps.mkString(", ")}" - ) + private[this] val init = Once { + val serviceLoadedTunableMap = LoadService[ServiceLoadedTunableMap]().toList + serviceLoadedTunableMap.foreach(setLoadedMap) + } + + private[this] def setLoadedMap(serviceLoadedTunableMap: ServiceLoadedTunableMap): Unit = { + loadedMaps.compute( + serviceLoadedTunableMap.id, + new BiFunction[String, TunableMap, TunableMap] { + def apply(id: String, curr: TunableMap): TunableMap = { + if (curr != null) + throw new IllegalStateException( + s"Found multiple `ServiceLoadedTunableMap`s for $id: $curr, $serviceLoadedTunableMap" + ) + else serviceLoadedTunableMap + } + } + ) + } + + /** + * Uses `com.twitter.app.LoadService` to load all [[ServiceLoadedTunableMap]]s, + * The subtype `com.twitter.util.tunable.ConfigbusTunableMap`s would re-construct + * and re-subscribe to the ConfigBus. + * + * @note this should be called after ServerInfo.initialized(). + */ + private[twitter] def reloadAll(): Unit = { + val serviceLoadedTunableMaps = LoadService[ServiceLoadedTunableMap]().toList + serviceLoadedTunableMaps.foreach { serviceLoadedTunableMap => + loadedMaps.computeIfPresent( + serviceLoadedTunableMap.id, + new BiFunction[String, TunableMap, TunableMap] { + def apply(id: String, curr: TunableMap): TunableMap = serviceLoadedTunableMap + } + ) } } } diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/Tunable.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/Tunable.scala index a65c4e0a88..088a319de2 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/Tunable.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/Tunable.scala @@ -1,5 +1,7 @@ package com.twitter.util.tunable +import com.twitter.util.Var + /** * A [[Tunable]] is an abstraction for an object that produces a Some(value) or None when applied. * Implementations may enable mutation, such that successive applications of the [[Tunable]] @@ -13,6 +15,12 @@ package com.twitter.util.tunable */ sealed abstract class Tunable[T](val id: String) { self => + /** + * Returns a [[Var]] that observes this [[Tunable]]. Observers of the [[Var]] will be notified + * whenever the value of this [[Tunable]] is updated. + */ + def asVar: Var[Option[T]] + // validate id is not empty if (id.trim.isEmpty) throw new IllegalArgumentException("Tunable id must not be empty") @@ -35,6 +43,12 @@ sealed abstract class Tunable[T](val id: String) { self => */ def orElse(that: Tunable[T]): Tunable[T] = new Tunable[T](id) { + + def asVar: Var[Option[T]] = self.asVar.flatMap { + case Some(_) => self.asVar + case None => that.asVar + } + override def toString: String = s"${self.toString}.orElse(${that.toString})" @@ -49,6 +63,9 @@ sealed abstract class Tunable[T](val id: String) { self => */ def map[A](f: T => A): Tunable[A] = new Tunable[A](id) { + + def asVar: Var[Option[A]] = self.asVar.map(_.map(f)) + def apply(): Option[A] = self().map(f) } } @@ -58,16 +75,17 @@ object Tunable { /** * A [[Tunable]] that always returns `value` when applied. */ - final class Const[T](id: String, private val value: T) extends Tunable[T](id) { - private[this] val SomeValue = Some(value) + final class Const[T](id: String, value: T) extends Tunable[T](id) { + private val holder = Var.value(Some(value)) + + def apply(): Option[T] = holder() - def apply(): Option[T] = - SomeValue + def asVar: Var[Option[T]] = holder } object Const { def unapply[T](tunable: Tunable[T]): Option[T] = tunable match { - case tunable: Tunable.Const[T] => Some(tunable.value) + case tunable: Tunable.Const[T] => tunable.holder() case _ => None } } @@ -78,8 +96,11 @@ object Tunable { def const[T](id: String, value: T): Const[T] = new Const(id, value) private[this] val NoneTunable = new Tunable[Any]("com.twitter.util.tunable.NoneTunable") { + def apply(): Option[Any] = None + + val asVar: Var[Option[Any]] = Var.value(None) } /** @@ -91,8 +112,11 @@ object Tunable { /** * A [[Tunable]] whose value can be changed. Operations are thread-safe. */ - final class Mutable[T] private[tunable] (id: String, @volatile private var _value: Option[T]) - extends Tunable[T](id) { + final class Mutable[T] private[tunable] (id: String, _value: Option[T]) extends Tunable[T](id) { + + private[this] final val mutableHolder = Var(_value) + + def asVar: Var[Option[T]] = mutableHolder /** * Set the value of the [[Tunable.Mutable]]. @@ -100,21 +124,18 @@ object Tunable { * Note that setting the value to `null` will result in a value of Some(null) when the * [[Tunable]] is applied. */ - def set(value: T): Unit = - _value = Some(value) + def set(value: T): Unit = mutableHolder.update(Some(value)) /** * Clear the value of the [[Tunable.Mutable]]. Calling `apply` on the [[Tunable.Mutable]] * will produce `None`. */ - def clear(): Unit = - _value = None + def clear(): Unit = mutableHolder.update(None) /** * Get the current value of the [[Tunable.Mutable]] */ - def apply(): Option[T] = - _value + def apply(): Option[T] = mutableHolder() } /** diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/TunableMap.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/TunableMap.scala index c9b31c0be3..8ac4b9aaaa 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/TunableMap.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/TunableMap.scala @@ -3,8 +3,8 @@ package com.twitter.util.tunable import java.util.concurrent.ConcurrentHashMap import java.util.function.BiFunction import scala.annotation.varargs -import scala.collection.JavaConverters._ import scala.collection.mutable +import scala.jdk.CollectionConverters._ /** * A Map that can be used to access [[Tunable]]s using [[TunableMap.Key]]s. @@ -13,7 +13,7 @@ abstract class TunableMap { self => /** * Returns a [[Tunable]] for the given `key.id` in the [[TunableMap]]. If the [[Tunable]] is not - * of type `T` or a subclass of type `T`, throws a [[ClassCastException]] + * of type `T` or a subclass of type `T`, throws a `ClassCastException` */ def apply[T](key: TunableMap.Key[T]): Tunable[T] @@ -40,8 +40,8 @@ abstract class TunableMap { self => def entries: Iterator[TunableMap.Entry[_]] /** - * Compose this [[TunableMap]] with another [[TunableMap]]. [[Tunables]] retrieved from - * the composed map prioritize the values of [[Tunables]] in the this map over the other + * Compose this [[TunableMap]] with another [[TunableMap]]. [[Tunable]]s retrieved from + * the composed map prioritize the values of [[Tunable]]s in the this map over the other * [[TunableMap]]. */ def orElse(that: TunableMap): TunableMap = (self, that) match { @@ -141,7 +141,7 @@ object TunableMap { * A [[TunableMap]] that can be updated via `put` and `clear` operations. Putting * a value for a given `id` will update the current value for the `id`, or create * a new [[Tunable]] if it does not exist. The type of the new value must match that of the - * existing value, or a [[ClassCastException]] will be thrown. + * existing value, or a `ClassCastException` will be thrown. * * `apply` returns a [[Tunable]] for a given [[TunableMap.Key]] and creates it if does not already * exist. Updates to the [[TunableMap]] update the underlying [[Tunable]]s; for example, a @@ -163,9 +163,7 @@ object TunableMap { private[tunable] def replace(newMap: TunableMap): Unit = { // clear keys no longer present - (entries.map(_.key).toSet -- newMap.entries.map(_.key).toSet).foreach { key => - clear(key) - } + (entries.map(_.key).toSet -- newMap.entries.map(_.key).toSet).foreach { key => clear(key) } // add/update new tunable values ++=(newMap) diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/linter/BUILD b/util-tunable/src/main/scala/com/twitter/util/tunable/linter/BUILD index 208d12477e..72c37656a6 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/linter/BUILD +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/linter/BUILD @@ -1,20 +1,24 @@ scala_library( name = "scala", - sources = rglobs("*.scala"), + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], dependencies = [ "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", "util/util-app/src/main/scala", - "util/util-core/src/main/scala", - "util/util-tunable/src/main/scala", + "util/util-core:util-core-util", + "util/util-tunable/src/main/scala/com/twitter/util/tunable", ], ) jvm_binary( name = "configuration-linter", main = "com.twitter.util.tunable.linter.ConfigurationLinter", - fatal_warnings = True, + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + tags = ["bazel-compatible"], dependencies = [ ":scala", ], diff --git a/util-tunable/src/main/scala/com/twitter/util/tunable/linter/ConfigurationLinter.scala b/util-tunable/src/main/scala/com/twitter/util/tunable/linter/ConfigurationLinter.scala index 13c2eb975c..4818004d85 100644 --- a/util-tunable/src/main/scala/com/twitter/util/tunable/linter/ConfigurationLinter.scala +++ b/util-tunable/src/main/scala/com/twitter/util/tunable/linter/ConfigurationLinter.scala @@ -3,7 +3,8 @@ package com.twitter.util.tunable.linter import com.twitter.app.App import com.twitter.util.{Throw, Return} import com.twitter.util.tunable.JsonTunableMapper -import java.net.URL + +import java.nio.file.Paths object ConfigurationLinter extends App { @@ -15,8 +16,7 @@ object ConfigurationLinter extends App { args.foreach { path => println(s"File: $path\n") - - val url = new URL("file://" + path) + val url = Paths.get(path).toUri().toURL() JsonTunableMapper().parse(url) match { case Return(map) => diff --git a/util-tunable/src/test/java/com/twitter/util/tunable/BUILD.bazel b/util-tunable/src/test/java/com/twitter/util/tunable/BUILD.bazel new file mode 100644 index 0000000000..629a404f97 --- /dev/null +++ b/util-tunable/src/test/java/com/twitter/util/tunable/BUILD.bazel @@ -0,0 +1,16 @@ +junit_tests( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scala-lang:scala-library", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-tunable/src/main/scala/com/twitter/util/tunable", + ], +) diff --git a/util-tunable/src/test/resources/BUILD b/util-tunable/src/test/resources/BUILD deleted file mode 100644 index cc44b5fc88..0000000000 --- a/util-tunable/src/test/resources/BUILD +++ /dev/null @@ -1,3 +0,0 @@ -resources( - sources = rglobs("*"), -) diff --git a/util-tunable/src/test/resources/BUILD.bazel b/util-tunable/src/test/resources/BUILD.bazel new file mode 100644 index 0000000000..301b666c21 --- /dev/null +++ b/util-tunable/src/test/resources/BUILD.bazel @@ -0,0 +1,4 @@ +resources( + sources = ["**/*"], + tags = ["bazel-compatible"], +) diff --git a/util-tunable/src/test/scala/BUILD b/util-tunable/src/test/scala/BUILD deleted file mode 100644 index c1409da3c6..0000000000 --- a/util-tunable/src/test/scala/BUILD +++ /dev/null @@ -1,15 +0,0 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, - dependencies = [ - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", - "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", - "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", - "3rdparty/jvm/junit", - "3rdparty/jvm/org/scalacheck", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-tunable/src/main/scala", - "util/util-tunable/src/test/resources", - ], -) diff --git a/util-tunable/src/test/scala/BUILD.bazel b/util-tunable/src/test/scala/BUILD.bazel new file mode 100644 index 0000000000..35a01a818a --- /dev/null +++ b/util-tunable/src/test/scala/BUILD.bazel @@ -0,0 +1,6 @@ +test_suite( + tags = ["bazel-compatible"], + dependencies = [ + "util/util-tunable/src/test/scala/com/twitter/util/tunable", + ], +) diff --git a/util-tunable/src/test/scala/com/twitter/util/tunable/BUILD.bazel b/util-tunable/src/test/scala/com/twitter/util/tunable/BUILD.bazel new file mode 100644 index 0000000000..7c0e487859 --- /dev/null +++ b/util-tunable/src/test/scala/com/twitter/util/tunable/BUILD.bazel @@ -0,0 +1,23 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-core", + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-databind", + "3rdparty/jvm/com/fasterxml/jackson/module:jackson-module-scala", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-tunable/src/main/scala/com/twitter/util/tunable", + "util/util-tunable/src/test/resources", + ], +) diff --git a/util-tunable/src/test/scala/com/twitter/util/tunable/JsonTunableMapperTest.scala b/util-tunable/src/test/scala/com/twitter/util/tunable/JsonTunableMapperTest.scala index 5d51c858c7..759d7c4a25 100644 --- a/util-tunable/src/test/scala/com/twitter/util/tunable/JsonTunableMapperTest.scala +++ b/util-tunable/src/test/scala/com/twitter/util/tunable/JsonTunableMapperTest.scala @@ -3,16 +3,19 @@ package com.twitter.util.tunable import com.fasterxml.jackson.core.JsonParser import com.fasterxml.jackson.databind.DeserializationContext import com.fasterxml.jackson.databind.deser.std.StdDeserializer -import com.twitter.conversions.storage._ -import com.twitter.conversions.time._ -import com.twitter.util.{Duration, Return, StorageUnit, Throw} -import org.scalatest.FunSuite -import scala.collection.JavaConverters._ +import com.twitter.conversions.StorageUnitOps._ +import com.twitter.conversions.DurationOps._ +import com.twitter.util.Duration +import com.twitter.util.Return +import com.twitter.util.StorageUnit +import com.twitter.util.Throw +import scala.jdk.CollectionConverters._ +import org.scalatest.funsuite.AnyFunSuite -// Used for veryifying custom deserialization +// Used for verifying custom deserialization case class Foo(number: Double) -class JsonTunableMapperTest extends FunSuite { +class JsonTunableMapperTest extends AnyFunSuite { test("returns a Throw if json is empty") { JsonTunableMapper().parse("") match { @@ -82,7 +85,8 @@ class JsonTunableMapperTest extends FunSuite { |{ "tunables": [ | { "id" : "timeoutId1", | "value" : "5.seconds", - | "type" : "com.twitter.util.Duration" + | "type" : "com.twitter.util.Duration", + | "comment": "a very important timeout" | }, | { "id" : "timeoutId2", | "value" : "Duration.Top", @@ -94,7 +98,8 @@ class JsonTunableMapperTest extends FunSuite { | }, | { "id" : "timeoutId4", | "value" : "Duration.Undefined", - | "type" : "com.twitter.util.Duration" + | "type" : "com.twitter.util.Duration", + | "comment": "You'll never believe it: *another* timeout." | } | ] |}""".stripMargin @@ -244,7 +249,8 @@ class JsonTunableMapperTest extends FunSuite { assert(map(TunableMap.Key[Duration]("timeoutId4"))() == Some(Duration.Undefined)) } - test("tunableMapForResources throws an Illegal argument exception when there are multiple paths") { + test( + "tunableMapForResources throws an Illegal argument exception when there are multiple paths") { val rsc = getClass.getClassLoader .getResources("com/twitter/tunables/IdForValidJson/instances.json") .asScala diff --git a/util-tunable/src/test/scala/com/twitter/util/tunable/ServiceLoadedTunableMapTest.scala b/util-tunable/src/test/scala/com/twitter/util/tunable/ServiceLoadedTunableMapTest.scala index 47b150d8c9..ddba96abb2 100644 --- a/util-tunable/src/test/scala/com/twitter/util/tunable/ServiceLoadedTunableMapTest.scala +++ b/util-tunable/src/test/scala/com/twitter/util/tunable/ServiceLoadedTunableMapTest.scala @@ -1,6 +1,6 @@ package com.twitter.util.tunable -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite class ServiceLoadedTunableTestClient1 extends ServiceLoadedTunableMap with TunableMap.Proxy { @@ -23,28 +23,34 @@ class ServiceLoadedTunableTestClient2Dup extends ServiceLoadedTunableMap with Tu def id: String = "client2" } -class ServiceLoadedTunableMapTest extends FunSuite { - - test("NullTunableMap returned when no matches") { - val tunableMap = ServiceLoadedTunableMap("Non-existent-id") - assert(tunableMap eq NullTunableMap) - } - - test("TunableMap returned when there is one match for id") { - val tunableMap = ServiceLoadedTunableMap("client1") - - assert(tunableMap.entries.size == 2) - assert(tunableMap(TunableMap.Key[String]("tunableId1"))() == Some("foo")) - assert(tunableMap(TunableMap.Key[Int]("tunableId2"))() == Some(5)) - } +class ServiceLoadedTunableMapTest extends AnyFunSuite { test( "IllegalArgumentException thrown when there is more than one ServiceLoadedTunableMap " + "for a given serviceName/id" ) { - intercept[IllegalStateException] { + val ex = intercept[IllegalStateException] { ServiceLoadedTunableMap("client2") } + assert(ex.getMessage.contains("Found multiple `ServiceLoadedTunableMap`s for client2")) + } + + test("NullTunableMap returned when no matches") { + intercept[IllegalStateException] { + val tunableMap = ServiceLoadedTunableMap("Non-existent-id") + assert(tunableMap eq NullTunableMap) + } + } + + test("TunableMap returned when there is one match for id") { + intercept[IllegalStateException] { + + val tunableMap = ServiceLoadedTunableMap("client1") + + assert(tunableMap.entries.size == 2) + assert(tunableMap(TunableMap.Key[String]("tunableId1"))() == Some("foo")) + assert(tunableMap(TunableMap.Key[Int]("tunableId2"))() == Some(5)) + } } } diff --git a/util-tunable/src/test/scala/com/twitter/util/tunable/TunableMapTest.scala b/util-tunable/src/test/scala/com/twitter/util/tunable/TunableMapTest.scala index 6b819a96c0..12a8c52597 100644 --- a/util-tunable/src/test/scala/com/twitter/util/tunable/TunableMapTest.scala +++ b/util-tunable/src/test/scala/com/twitter/util/tunable/TunableMapTest.scala @@ -1,8 +1,8 @@ package com.twitter.util.tunable -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class MutableTest extends FunSuite { +class MutableTest extends AnyFunSuite { test("map is initially empty") { assert(TunableMap.newMutable().entries.size == 0) @@ -336,9 +336,7 @@ class MutableTest extends FunSuite { map1.put("key2", 5) assert(combined.entries.size == 1) - assert(combined.entries.exists { entry => - entry.key.id == "key2" && entry.value == 5 - }) + assert(combined.entries.exists { entry => entry.key.id == "key2" && entry.value == 5 }) } test("orElse: entries returns entries from both maps") { @@ -353,12 +351,8 @@ class MutableTest extends FunSuite { map2.put("key2", 5) assert(combined.entries.size == 2) - assert(combined.entries.exists { entry => - entry.key.id == "key1" && entry.value == "hello1" - }) - assert(combined.entries.exists { entry => - entry.key.id == "key2" && entry.value == 5 - }) + assert(combined.entries.exists { entry => entry.key.id == "key1" && entry.value == "hello1" }) + assert(combined.entries.exists { entry => entry.key.id == "key2" && entry.value == 5 }) } test("orElse on two NullTunableMaps produces NullTunableMap") { @@ -390,7 +384,7 @@ class MutableTest extends FunSuite { } } -class NullTunableMapTest extends FunSuite { +class NullTunableMapTest extends AnyFunSuite { test("NullTunableMap returns Tunable.None") { val nullTunableMap = NullTunableMap diff --git a/util-tunable/src/test/scala/com/twitter/util/tunable/TunableTest.scala b/util-tunable/src/test/scala/com/twitter/util/tunable/TunableTest.scala index 72cb1d3908..0a7c7ed32c 100644 --- a/util-tunable/src/test/scala/com/twitter/util/tunable/TunableTest.scala +++ b/util-tunable/src/test/scala/com/twitter/util/tunable/TunableTest.scala @@ -1,8 +1,8 @@ package com.twitter.util.tunable -import org.scalatest.FunSuite +import org.scalatest.funsuite.AnyFunSuite -class TunableTest extends FunSuite { +class TunableTest extends AnyFunSuite { test("Tunable cannot have empty id") { intercept[IllegalArgumentException] { @@ -170,4 +170,67 @@ class TunableTest extends FunSuite { assert(composed() == None) } + + test("mutable tunables are observable") { + val tunable = Tunable.mutable("id", 5) + val obs = tunable.asVar + + assert(obs.sample().contains(5)) + } + + test("empty mutable tunables are observable") { + val tunable = Tunable.emptyMutable[Int]("id") + val obs = tunable.asVar + + assert(obs.sample().isEmpty) + tunable.set(5) + assert(obs.sample().contains(5)) + } + + test("mutable tunable changes are observed") { + val tunable = Tunable.mutable("id", 5) + val obs = tunable.asVar + + assert(obs.sample().contains(5)) + tunable.set(6) + assert(obs.sample().contains(6)) + } + + test("orElse Tunable is observable if first value is defined") { + val tunable = Tunable.mutable("a", 5) + val orElsed = tunable.orElse(Tunable.const("b", 6)) + + val obs = orElsed.asVar + assert(obs.sample().contains(5)) + tunable.clear() + assert(obs.sample().contains(6)) + } + + test("orElse Tunable is observable if first value is undefined") { + val tunable = Tunable.emptyMutable[Int]("a") + val orElsed = tunable.orElse(Tunable.const("b", 5)) + + val obs = orElsed.asVar + assert(obs.sample().contains(5)) + tunable.set(6) + assert(obs.sample().contains(6)) + } + + test("map of mutable Tunables is observable") { + val tunable = Tunable.mutable("id", 5) + val mapped = tunable.map(_ + 1) + val obs = mapped.asVar + + assert(obs.sample().contains(6)) + tunable.set(6) + assert(obs.sample().contains(7)) + } + + test("map of const Tunables is observable") { + val tunable = Tunable.const("id", 5) + val mapped = tunable.map(_ + 1) + val obs = mapped.asVar + + assert(obs.sample().contains(6)) + } } diff --git a/util-validator-constraints/PROJECT b/util-validator-constraints/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-validator-constraints/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/BUILD b/util-validator-constraints/src/main/java/com/twitter/util/validation/BUILD new file mode 100644 index 0000000000..4a6f8e42a5 --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/BUILD @@ -0,0 +1,18 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-java", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], +) diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/MethodValidation.java b/util-validator-constraints/src/main/java/com/twitter/util/validation/MethodValidation.java new file mode 100644 index 0000000000..9aee33853c --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/MethodValidation.java @@ -0,0 +1,47 @@ +package com.twitter.util.validation; + +import java.lang.annotation.Documented; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +/** + * A type of [[jakarta.validation.Constraint]] used to annotate case class methods + * to invoke for additional instance state validation. + */ +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +public @interface MethodValidation { + + /** Every constraint annotation must define a message element of type String. */ + String message() default ""; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. * + */ + Class[] payload() default {}; + + /** + * Specify the fields validated by the method. Used in error reporting. + */ + String[] fields() default {}; +} diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/BUILD b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/BUILD new file mode 100644 index 0000000000..c2610076fe --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/BUILD @@ -0,0 +1,18 @@ +java_library( + sources = ["*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-java-constraints", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], +) diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/CountryCode.java b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/CountryCode.java new file mode 100644 index 0000000000..9a8f5ba719 --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/CountryCode.java @@ -0,0 +1,52 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Repeatable; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import com.twitter.util.validation.constraints.CountryCode.List; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +@Repeatable(List.class) +public @interface CountryCode { + + /** Every constraint annotation must define a message element of type String. */ + String message() default "{com.twitter.util.validation.constraints.CountryCode.message}"; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. + */ + Class[] payload() default {}; + + /** + * Defines several {@code @CountryCode} annotations on the same element. + */ + @Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) + @Retention(RUNTIME) + @Documented + public @interface List { + CountryCode[] value(); + } +} diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/OneOf.java b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/OneOf.java new file mode 100644 index 0000000000..275d148653 --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/OneOf.java @@ -0,0 +1,57 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Repeatable; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import com.twitter.util.validation.constraints.OneOf.List; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +@Repeatable(List.class) +public @interface OneOf { + + /** Every constraint annotation must define a message element of type String. */ + String message() default "{com.twitter.util.validation.constraints.OneOf.message}"; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. * + */ + Class[] payload() default {}; + + /** + * The array of valid values. + */ + String[] value(); + + /** + * Defines several {@code @OneOf} annotations on the same element. + */ + @Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) + @Retention(RUNTIME) + @Documented + public @interface List { + OneOf[] value(); + } +} diff --git a/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/UUID.java b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/UUID.java new file mode 100644 index 0000000000..7a5b969db3 --- /dev/null +++ b/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints/UUID.java @@ -0,0 +1,52 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Repeatable; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import com.twitter.util.validation.constraints.UUID.List; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +@Repeatable(List.class) +public @interface UUID { + + /** Every constraint annotation must define a message element of type String. */ + String message() default "{com.twitter.util.validation.constraints.UUID.message}"; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. * + */ + Class[] payload() default {}; + + /** + * Defines several {@code @UUID} annotations on the same element. + */ + @Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) + @Retention(RUNTIME) + @Documented + public @interface List { + UUID[] value(); + } +} diff --git a/util-validator-constraints/src/main/resources/BUILD b/util-validator-constraints/src/main/resources/BUILD new file mode 100644 index 0000000000..ad57fee95f --- /dev/null +++ b/util-validator-constraints/src/main/resources/BUILD @@ -0,0 +1,6 @@ +resources( + sources = [ + "**/*.properties", + ], + tags = ["bazel-compatible"], +) diff --git a/util-validator-constraints/src/main/resources/ValidationMessages.properties b/util-validator-constraints/src/main/resources/ValidationMessages.properties new file mode 100644 index 0000000000..30c2277f07 --- /dev/null +++ b/util-validator-constraints/src/main/resources/ValidationMessages.properties @@ -0,0 +1,3 @@ +com.twitter.util.validation.constraints.CountryCode.message=${validatedValue} not a valid country code +com.twitter.util.validation.constraints.OneOf.message=${validatedValue} not one of {value} +com.twitter.util.validation.constraints.UUID.message=must be a valid UUID diff --git a/util-validator/PROJECT b/util-validator/PROJECT new file mode 100644 index 0000000000..cda5765dfa --- /dev/null +++ b/util-validator/PROJECT @@ -0,0 +1,7 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - csl-team@twitter.com diff --git a/util-validator/src/main/resources/BUILD b/util-validator/src/main/resources/BUILD new file mode 100644 index 0000000000..a59b2ee358 --- /dev/null +++ b/util-validator/src/main/resources/BUILD @@ -0,0 +1,6 @@ +resources( + sources = [ + "**/*.xml", + ], + tags = ["bazel-compatible"], +) diff --git a/util-validator/src/main/resources/META-INF/validation.xml b/util-validator/src/main/resources/META-INF/validation.xml new file mode 100644 index 0000000000..a2678728f0 --- /dev/null +++ b/util-validator/src/main/resources/META-INF/validation.xml @@ -0,0 +1,8 @@ + + com/twitter/util/validation/constraint-mappings.xml + diff --git a/util-validator/src/main/resources/com/twitter/util/validation/constraint-mappings.xml b/util-validator/src/main/resources/com/twitter/util/validation/constraint-mappings.xml new file mode 100644 index 0000000000..2622a2bf04 --- /dev/null +++ b/util-validator/src/main/resources/com/twitter/util/validation/constraint-mappings.xml @@ -0,0 +1,32 @@ + + + + com.twitter.util.validation.internal.validators.ISO3166CountryCodeConstraintValidator + + + + + com.twitter.util.validation.internal.validators.OneOfConstraintValidator + + + + + com.twitter.util.validation.internal.validators.UUIDConstraintValidator + + + + + com.twitter.util.validation.internal.validators.notempty.NotEmptyValidatorForIterable + + + + + com.twitter.util.validation.internal.validators.size.SizeValidatorForIterable + + + diff --git a/util-validator/src/main/scala/com/twitter/util/validation/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/BUILD new file mode 100644 index 0000000000..b715b02ab6 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/BUILD @@ -0,0 +1,52 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-core", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-core:scala", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints", + "util/util-validator-constraints/src/main/resources", # loads validation messages + "util/util-validator/src/main/resources", # loads custom constraint mappings + "util/util-validator/src/main/scala/com/twitter/util/validation/cfg", + "util/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation", + "util/util-validator/src/main/scala/com/twitter/util/validation/conversions", + "util/util-validator/src/main/scala/com/twitter/util/validation/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/executable", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/validators", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + "util/util-validator/src/main/scala/org/hibernate/validator/internal/engine", + "util/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints", + "util/util-validator/src/main/scala/com/twitter/util/validation/cfg", + "util/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation", + "util/util-validator/src/main/scala/com/twitter/util/validation/conversions", + "util/util-validator/src/main/scala/com/twitter/util/validation/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/executable", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/DescriptorFactory.scala b/util-validator/src/main/scala/com/twitter/util/validation/DescriptorFactory.scala new file mode 100644 index 0000000000..4d4818a0dc --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/DescriptorFactory.scala @@ -0,0 +1,517 @@ +package com.twitter.util.validation + +import com.github.benmanes.caffeine.cache.{Cache, Caffeine} +import com.twitter.util.reflect.{Annotations, Types => ReflectTypes} +import com.twitter.util.validation.engine.MethodValidationResult +import com.twitter.util.validation.internal.Types +import com.twitter.util.validation.metadata.{ + CaseClassDescriptor, + ConstructorDescriptor, + ExecutableDescriptor, + MethodDescriptor, + PropertyDescriptor +} +import com.twitter.util.{Return, Try} +import jakarta.validation.{Constraint, ConstraintDeclarationException, Valid, ValidationException} +import java.lang.annotation.Annotation +import java.lang.reflect.{Constructor, Executable, Method, Parameter} +import org.hibernate.validator.internal.metadata.core.ConstraintHelper +import org.json4s.reflect.{ClassDescriptor, ConstructorParamDescriptor, Reflector, ScalaType} +import org.json4s.{reflect => json4s} +import scala.collection.mutable +import scala.util.control.NonFatal + +private[validation] object DescriptorFactory { + + def unmangleName(name: String): String = scala.reflect.NameTransformer.decode(name) + def mangleName(name: String): String = scala.reflect.NameTransformer.encode(name) + + /** If the given class is marked for cascaded validation or not. True if the class is a case class and has the `@Valid` annotation */ + private def isCascadedValidation(erasure: Class[_], annotations: Array[Annotation]): Boolean = + Annotations.findAnnotation[Valid](annotations).isDefined && ReflectTypes.isCaseClass(erasure) + + private case class ConstructorParams( + param: ConstructorParamDescriptor, + annotations: Array[Annotation]) + + private def getCascadedScalaType( + scalaType: ScalaType, + annotations: Array[Annotation] + ): Option[ScalaType] = + Types.getContainedScalaType(scalaType) match { + case Some(containedScalaType) + if isCascadedValidation(containedScalaType.erasure, annotations) => + Some(containedScalaType) + case _ => None + } +} + +/** + * Used to describe a given Class[T] as a CaseClassDescriptor[T] + */ +private[validation] class DescriptorFactory( + descriptorCacheSize: Long, + constraintHelper: ConstraintHelper) { + + import DescriptorFactory._ + + // A caffeine cache to store the expensive reflection calls on the same object. Caffeine cache + // uses the `Window TinyLfu` policy to remove evicted keys. + // For more information, check out: https://github.com/ben-manes/caffeine/wiki/Efficiency + private[validation] val caseClassDescriptors: Cache[Class[_], CaseClassDescriptor[_]] = + Caffeine + .newBuilder() + .maximumSize(descriptorCacheSize) + .build[Class[_], CaseClassDescriptor[_]]() + + /** Close and attempt to clean up resources */ + def close(): Unit = { + caseClassDescriptors.invalidateAll() + caseClassDescriptors.cleanUp() + } + + /** + * Describe a [[Class]]. + * + * @note the returned [[CaseClassDescriptor]] is cached for repeated lookup attempts keyed by + * the given Class[T] type. + */ + def describe[T]( + clazz: Class[T] + ): CaseClassDescriptor[T] = { + caseClassDescriptors + .get( + clazz, + (key: Class[_]) => buildDescriptor[T](clazz) + ).asInstanceOf[CaseClassDescriptor[T]] + } + + /** + * Describe a [[Constructor]]. + * + * @note the returned [[ConstructorDescriptor]] is NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + def describe[T]( + constructor: Constructor[T] + ): ConstructorDescriptor = { + buildConstructorDescriptor( + constructor.getDeclaringClass, + getJson4sConstructorDescriptor(constructor), + None + ) + } + + /** + * Describe a `@MethodValidation`-annotated or otherwise constrained [[Method]]. + * + * As we do not want to describe every possible case class method, this potentially + * returns a None in the case where the method has no constraint annotation and no + * constrained parameters. + * + * @note the returned [[MethodDescriptor]] is NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + def describe(method: Method): Option[MethodDescriptor] = buildMethodDescriptor(method) + + /** + * Describe an [[Executable]] given an optional "mix-in" Class. + * + * @note the returned [[ExecutableDescriptor]] is NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + def describeExecutable[T]( + executable: Executable, + mixinClazz: Option[Class[_]] + ): ExecutableDescriptor = { + executable match { + case constructor: Constructor[_] => + val desc = getJson4sConstructorDescriptor(constructor) + val clazz = constructor.getDeclaringClass.asInstanceOf[Class[T]] + val parameterAnnotations = mixinClazz match { + case Some(mixin) => + val constructorAnnotationMap = getConstructorParams(clazz, desc) + // augment with mixin class field annotations + constructorAnnotationMap.map { + case (name, params) => + try { + val method = mixin.getDeclaredMethod(name) + ( + name, + params.copy(annotations = + params.annotations ++ + method.getAnnotations.filter(isConstraintAnnotation)) + ) + } catch { + case NonFatal(_) => // do nothing + (name, params) + } + } + case _ => + getConstructorParams(clazz, desc) + } + buildConstructorDescriptor( + constructor.getDeclaringClass, + desc, + Some(parameterAnnotations) + ) + case method: Method => + MethodDescriptor( + method = method, + annotations = method.getDeclaredAnnotations.filter(isConstraintAnnotation), + members = method.getParameters.map { parameter => + val name = parameter.getName + // augment with mixin class field annotations + val parameterAnnotations = + mixinClazz match { + case Some(mixin) => + try { + val method = mixin.getDeclaredMethod(name) + parameter.getAnnotations ++ + method.getAnnotations.filter(isConstraintAnnotation) + } catch { + case NonFatal(_) => // do nothing + parameter.getAnnotations + } + case _ => // do nothing + parameter.getAnnotations + } + parameter.getName -> buildPropertyDescriptor( + Reflector.scalaTypeOf(parameter.getParameterizedType), + parameterAnnotations + ) + }.toMap + ) + case _ => // should not get here + throw new IllegalArgumentException + } + } + + /** + * Describe `@MethodValidation`-annotated or otherwise constrained methods of a given `Class[_]`. + * + * @note the returned [[MethodDescriptor]] instances are NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + def describeMethods(clazz: Class[_]): Array[MethodDescriptor] = { + val clazzMethods = clazz.getMethods + if (clazzMethods.nonEmpty) { + val methods = new mutable.ArrayBuffer[MethodDescriptor](clazzMethods.length) + var index = 0 + val length = clazzMethods.length + while (index < length) { + val method = clazzMethods(index) + buildMethodDescriptor(method).foreach(methods.append(_)) + index += 1 + } + methods.toArray + } else Array.empty[MethodDescriptor] + } + + /* Private */ + + private[this] def buildDescriptor[T](clazz: Class[T]): CaseClassDescriptor[T] = { + // for case classes, annotations for params only appear on the executable + val clazzDescriptor: ClassDescriptor = Reflector.describe(clazz).asInstanceOf[ClassDescriptor] + // map of parameter name to a class with all annotations for the parameter (including inherited) + val constructorParams: scala.collection.Map[String, ConstructorParams] = + getConstructorParams(clazz, clazzDescriptor) + + val members = new mutable.HashMap[String, PropertyDescriptor]() + // we validate only declared properties of the class + val properties: Seq[json4s.PropertyDescriptor] = clazzDescriptor.properties + var index = 0 + val size = properties.size + while (index < size) { + val property = properties(index) + val (fromConstructorScalaType, annotations) = { + constructorParams.get(property.mangledName) match { + case Some(constructorParam) => + // we already have information for the field as it is an annotated executable parameter + (constructorParam.param.argType, constructorParam.annotations) + case _ => + // we don't have information, use the property information + (property.returnType, property.field.getAnnotations) + } + } + + // create property descriptorFactory + members.put( + property.mangledName, + buildPropertyDescriptor(fromConstructorScalaType, annotations) + ) + index += 1 + } + + val constructors = new mutable.ArrayBuffer[ConstructorDescriptor]() + // create executable descriptors + var jindex = 0 + val jsize = clazzDescriptor.constructors.size + while (jindex < jsize) { + constructors.append( + buildConstructorDescriptor(clazz, clazzDescriptor.constructors(jindex), None) + ) + jindex += 1 + } + + val methods = new mutable.ArrayBuffer[MethodDescriptor]() + val clazzMethods = clazz.getMethods + if (clazzMethods.nonEmpty) { + var index = 0 + val length = clazzMethods.length + while (index < length) { + val method = clazzMethods(index) + buildMethodDescriptor(method).foreach(methods.append(_)) + index += 1 + } + } + + CaseClassDescriptor( + clazz = clazz, + scalaType = clazzDescriptor.erasure, + annotations = clazzDescriptor.erasure.erasure.getAnnotations.filter(isConstraintAnnotation), + constructors = constructors.toArray, + members = members.toMap, + methods = methods.toArray + ) + } + + /* Exposed for testing */ + // find the executable with parameters that are all declared class fields + private[validation] def findConstructor( + clazzDescriptor: ClassDescriptor + ): org.json4s.reflect.ConstructorDescriptor = { + def isDefaultConstructorDescriptor( + constructorDescriptor: org.json4s.reflect.ConstructorDescriptor + ): Boolean = { + constructorDescriptor.isPrimary || constructorDescriptor.params.forall { param => + Try(clazzDescriptor.erasure.erasure.getDeclaredField(param.mangledName)).isReturn + } + } + + // locate the executable where all params are also declared fields else error + clazzDescriptor.constructors + .find(isDefaultConstructorDescriptor) + .getOrElse(throw new ValidationException( + s"Unable to parse case class for validation: ${clazzDescriptor.erasure.fullName}")) + } + + private[this] def getConstructorParams( + clazz: Class[_], + clazzDescriptor: ClassDescriptor + ): scala.collection.Map[String, ConstructorParams] = { + // find the default executable descriptorFactory (the underlying executable + // may be a executable or a method) + val constructorDescriptor: org.json4s.reflect.ConstructorDescriptor = + findConstructor(clazzDescriptor) + getConstructorParams( + clazz, + constructorDescriptor + ) + } + + private[this] def getConstructorParams( + clazz: Class[_], + constructorDescriptor: org.json4s.reflect.ConstructorDescriptor + ): scala.collection.Map[String, ConstructorParams] = { + val constructorParameters: Array[Parameter] = { + Option(constructorDescriptor.constructor.constructor) match { + case Some(cons) => + cons.getParameters + case _ => + // factory method "executable" + constructorDescriptor.constructor.method.getParameters + } + } + + // find all inherited annotations for every executable param + val allFieldAnnotations = + findAnnotations( + clazz, + clazz.getDeclaredFields.map(_.getName).toSet, + constructorDescriptor.params.map { param => + // annotations can be on the executable param OR they were + // copied to the generated field with a meta annotation. we + // rely on the fact that a given annotation should not show up in + // both arrays the way the Scala compiler currently works, however + // the executable we're iterating through may not represent + // actual declared fields in the class, so we have to apply logic + param.mangledName -> getAnnotations( + param.mangledName, + constructorParameters(param.argIndex), + clazz) + }.toMap + ) + + val result: mutable.HashMap[String, ConstructorParams] = + new mutable.HashMap[String, ConstructorParams]() + var index = 0 + val length = constructorParameters.length + while (index < length) { + val descriptor = constructorDescriptor.params(index) + val filteredAnnotations = allFieldAnnotations(descriptor.mangledName).filter { ann => + Annotations.isAnnotationPresent[Constraint](ann) || + Annotations.equals[Valid](ann) + } + + result.put( + descriptor.mangledName, + ConstructorParams( + descriptor, + filteredAnnotations + ) + ) + index += 1 + } + result + } + + private[this] def getAnnotations( + name: String, + parameter: Parameter, + clazz: Class[_] + ): Array[Annotation] = { + val fromClazzAnnotations: Array[Annotation] = + Try(clazz.getDeclaredField(name)) match { + case Return(declaredField) => + declaredField.getAnnotations + case _ => + Array.empty[Annotation] + } + parameter.getAnnotations ++ fromClazzAnnotations + } + + private[this] def findAnnotations( + clazz: Class[_], + declaredFields: Set[String], + fieldAnnotations: scala.collection.Map[String, Array[Annotation]] + ): scala.collection.Map[String, Array[Annotation]] = { + val collectorMap = new scala.collection.mutable.HashMap[String, Array[Annotation]]() + collectorMap ++= fieldAnnotations + // find inherited annotations + Annotations.findDeclaredAnnotations( + clazz, + declaredFields, + collectorMap + ) + } + + private[this] def getJson4sConstructorDescriptor( + constructor: Constructor[_] + ): org.json4s.reflect.ConstructorDescriptor = { + val parameters: Array[Parameter] = constructor.getParameters + org.json4s.reflect.ConstructorDescriptor( + params = Seq.tabulate(parameters.length) { index => + val parameter = parameters(index) + org.json4s.reflect.ConstructorParamDescriptor( + name = unmangleName(parameter.getName), + mangledName = parameter.getName, + argIndex = index, + argType = Reflector.scalaTypeOf(parameter.getParameterizedType), + defaultValue = None + ) + }, + constructor = new org.json4s.reflect.Executable(constructor, true), + isPrimary = true + ) + } + + private[this] def buildConstructorDescriptor[T]( + clazz: Class[T], + constructorDescriptor: org.json4s.reflect.ConstructorDescriptor, + parameterAnnotationsMap: Option[scala.collection.Map[String, ConstructorParams]] + ): ConstructorDescriptor = { + val parameterAnnotations = parameterAnnotationsMap match { + case Some(m) => m + case _ => getConstructorParams(clazz, constructorDescriptor) + } + val (executable, annotations) = Option(constructorDescriptor.constructor.constructor) match { + case Some(constructor) => + (constructor, constructor.getAnnotations) + case _ => + ( + constructorDescriptor.constructor.method, + constructorDescriptor.constructor.method.getAnnotations) + } + + ConstructorDescriptor( + executable, + annotations = annotations.filter(isConstraintAnnotation), + members = constructorDescriptor.params.map { parameter => + parameter.mangledName -> buildPropertyDescriptor( + parameter.argType, + parameterAnnotations(parameter.mangledName).annotations + ) + }.toMap + ) + } + + /** Create a PropertyDescriptor from the given ScalaType and annotations */ + private[this] def buildPropertyDescriptor[T]( + scalaType: ScalaType, + annotations: Array[Annotation] + ): PropertyDescriptor = { + val cascadedScalaType = getCascadedScalaType(scalaType, annotations) + val isCascaded = // annotated with @Valid and is a case class type + Annotations.findAnnotation(classOf[Valid], annotations).isDefined && + cascadedScalaType.exists(p => ReflectTypes.isCaseClass(p.erasure)) + + PropertyDescriptor( + scalaType = scalaType, + cascadedScalaType = cascadedScalaType, + annotations = annotations.filter(isConstraintAnnotation), + isCascaded = isCascaded + ) + } + + /** Create a MethodDescriptor for a given clazz Method */ + private[this] def buildMethodDescriptor(method: Method): Option[MethodDescriptor] = { + if (isMethodValidation(method) && checkMethodValidationMethod(method)) { + Some( + MethodDescriptor( + method = method, + annotations = method.getDeclaredAnnotations.filter(isConstraintAnnotation), + members = Map.empty + ) + ) + } else if (isConstrainedMethod(method) || hasConstrainedParameters(method)) { + Some( + MethodDescriptor( + method = method, + annotations = method.getDeclaredAnnotations.filter(isConstraintAnnotation), + members = method.getParameters.map { parameter => + parameter.getName -> buildPropertyDescriptor( + Reflector.scalaTypeOf(parameter.getParameterizedType), + parameter.getAnnotations + ) + }.toMap + ) + ) + } else None + } + + private[this] def checkMethodValidationMethod(method: Method): Boolean = { + if (method.getParameterCount != 0) + throw new ConstraintDeclarationException( + s"Methods annotated with @${classOf[MethodValidation].getSimpleName} must not declare any arguments") + if (method.getReturnType != classOf[MethodValidationResult]) + throw new ConstraintDeclarationException(s"Methods annotated with @${classOf[ + MethodValidation].getSimpleName} must return a ${classOf[MethodValidationResult].getName}") + true + } + + private[this] def isConstraintAnnotation(annotation: Annotation): Boolean = + constraintHelper.isConstraintAnnotation(annotation.annotationType()) + + // Array of annotations contains a constraint annotation + private[this] def isConstrainedMethod(method: Method): Boolean = + method.getDeclaredAnnotations.exists(isConstraintAnnotation) + + private[this] def hasConstrainedParameters(method: Method): Boolean = { + method.getParameters.exists(_.getAnnotations.exists(isConstraintAnnotation)) + } + + // Array of annotation contains @MethodValidation + private[this] def isMethodValidation(method: Method): Boolean = + Annotations.findAnnotation[MethodValidation](method.getDeclaredAnnotations).isDefined +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/ScalaValidator.scala b/util-validator/src/main/scala/com/twitter/util/validation/ScalaValidator.scala new file mode 100644 index 0000000000..c204ac2fc9 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/ScalaValidator.scala @@ -0,0 +1,1507 @@ +package com.twitter.util.validation + +import com.twitter.util.Memoize +import com.twitter.util.logging.Logging +import com.twitter.util.reflect.{Classes, Annotations => ReflectAnnotations, Types => ReflectTypes} +import com.twitter.util.validation.cfg.ConstraintMapping +import com.twitter.util.validation.engine.MethodValidationResult +import com.twitter.util.validation.executable.ScalaExecutableValidator +import com.twitter.util.validation.internal.constraintvalidation.ConstraintValidatorContextFactory +import com.twitter.util.validation.internal.engine.ConstraintViolationFactory +import com.twitter.util.validation.internal.metadata.descriptor.{ + ConstraintDescriptorFactory, + MethodValidationConstraintDescriptor +} +import com.twitter.util.validation.internal.{AnnotationFactory, Types, ValidationContext} +import com.twitter.util.validation.metadata._ +import jakarta.validation._ +import jakarta.validation.constraintvalidation.{SupportedValidationTarget, ValidationTarget} +import jakarta.validation.metadata.ConstraintDescriptor +import jakarta.validation.spi.ValidationProvider +import java.lang.annotation.Annotation +import java.lang.reflect.{Constructor, Executable, InvocationTargetException, Method, Parameter} +import java.util.Collections +import org.hibernate.validator.{HibernateValidator, HibernateValidatorConfiguration} +import org.hibernate.validator.internal.engine.constraintvalidation.ConstraintValidatorManager +import org.hibernate.validator.internal.engine.path.PathImpl +import org.hibernate.validator.internal.engine.{ValidatorFactoryImpl, ValidatorFactoryInspector} +import org.hibernate.validator.internal.metadata.aggregated.{ + CascadingMetaDataBuilder, + ExecutableMetaData +} +import org.hibernate.validator.internal.metadata.descriptor.ConstraintDescriptorImpl +import org.hibernate.validator.internal.metadata.raw.ConstrainedElement.ConstrainedElementKind +import org.hibernate.validator.internal.metadata.raw.{ + ConfigurationSource, + ConstrainedExecutable, + ConstrainedParameter +} +import org.hibernate.validator.internal.properties.javabean.{ + ConstructorCallable, + ExecutableCallable, + MethodCallable +} +import org.json4s.reflect.{Reflector, ScalaType} +import scala.annotation.{tailrec, varargs} +import scala.collection.mutable +import scala.collection.mutable.ListBuffer +import scala.jdk.CollectionConverters._ +import scala.util.control.NonFatal + +object ScalaValidator { + + /** The size of the caffeine cache that is used to store reflection data on a validated case class. */ + private val DefaultDescriptorCacheSize: Long = 128 + + case class Builder private[validation] ( + descriptorCacheSize: Long = DefaultDescriptorCacheSize, + messageInterpolator: Option[MessageInterpolator] = None, + constraintMappings: Set[ConstraintMapping] = Set.empty) { + + /** + * Define the size of the cache that stores annotation data for validated case classes. + */ + def withDescriptorCacheSize(size: Long): ScalaValidator.Builder = + Builder( + descriptorCacheSize = size, + messageInterpolator = this.messageInterpolator, + constraintMappings = this.constraintMappings + ) + + def withMessageInterpolator( + messageInterpolator: MessageInterpolator + ): ScalaValidator.Builder = + Builder( + descriptorCacheSize = this.descriptorCacheSize, + messageInterpolator = Some(messageInterpolator), + constraintMappings = this.constraintMappings + ) + + /** + * Register a Set of [[ConstraintMapping]]. + * + * @note Java users please see the version which takes a [[java.util.Set]] + */ + def withConstraintMappings( + constraintMappings: Set[ConstraintMapping] + ): ScalaValidator.Builder = + Builder( + descriptorCacheSize = this.descriptorCacheSize, + messageInterpolator = this.messageInterpolator, + constraintMappings = constraintMappings + ) + + /** + * Register a Set of [[ConstraintMapping]]. + * + * @note Scala users please see the version which takes a [[Set]] + */ + def withConstraintMappings( + constraintMappings: java.util.Set[ConstraintMapping] + ): ScalaValidator.Builder = + Builder( + descriptorCacheSize = this.descriptorCacheSize, + messageInterpolator = this.messageInterpolator, + constraintMappings = constraintMappings.asScala.toSet + ) + + /** + * Register a [[ConstraintMapping]]. + */ + def withConstraintMapping( + constraintMapping: ConstraintMapping + ): ScalaValidator.Builder = + Builder( + descriptorCacheSize = this.descriptorCacheSize, + messageInterpolator = this.messageInterpolator, + constraintMappings = Set(constraintMapping) + ) + + def validator: ScalaValidator = { + val configuration: HibernateValidatorConfiguration = + Validation + .byProvider[ + HibernateValidatorConfiguration, + ValidationProvider[HibernateValidatorConfiguration]]( + classOf[HibernateValidator] + .asInstanceOf[Class[ValidationProvider[HibernateValidatorConfiguration]]]) + .configure() + + // add user-configured message interpolator + messageInterpolator.map { interpolator => + configuration.messageInterpolator(interpolator) + } + + // add user-configured constraint mappings + constraintMappings.foreach { constraintMapping => + val hibernateConstraintMapping: org.hibernate.validator.cfg.ConstraintMapping = + configuration.createConstraintMapping() + hibernateConstraintMapping + .constraintDefinition(constraintMapping.annotationType.asInstanceOf[Class[Annotation]]) + .includeExistingValidators(constraintMapping.includeExistingValidators) + .validatedBy(constraintMapping.constraintValidator + .asInstanceOf[Class[_ <: ConstraintValidator[Annotation, _]]]) + configuration.addMapping(hibernateConstraintMapping) + } + + val validatorFactory = + configuration.buildValidatorFactory + + new ScalaValidator( + this.descriptorCacheSize, + new ValidatorFactoryInspector(validatorFactory.asInstanceOf[ValidatorFactoryImpl]), + validatorFactory.getValidator + ) + } + } + + def builder: ScalaValidator.Builder = Builder() + + def apply(): ScalaValidator = ScalaValidator.builder.validator +} + +class ScalaValidator private[validation] ( + cacheSize: Long, + validatorFactory: ValidatorFactoryInspector, + val underlying: Validator) + extends ScalaExecutableValidator + with Logging { + + /* Exposed for testing */ + private[validation] val descriptorFactory: DescriptorFactory = + new DescriptorFactory(cacheSize, validatorFactory.constraintHelper) + + private[this] val constraintViolationFactory: ConstraintViolationFactory = + new ConstraintViolationFactory(validatorFactory) + + private[this] val constraintDescriptorFactory: ConstraintDescriptorFactory = + new ConstraintDescriptorFactory(validatorFactory) + + private[this] val constraintValidatorContextFactory: ConstraintValidatorContextFactory = + new ConstraintValidatorContextFactory(validatorFactory) + + private[this] val constraintValidatorManager: ConstraintValidatorManager = + validatorFactory.getConstraintCreationContext.getConstraintValidatorManager + + /** Close resources.*/ + def close(): Unit = { + descriptorFactory.close() + validatorFactory.close() + } + + /** + * Returns a [[CaseClassDescriptor]] object describing the given case class type constraints. + * Descriptors describe constraints on a given class and any cascaded case class types. + * + * The returned object (and associated objects including CaseClassDescriptors) are immutable. + * + * @param clazz class or interface type evaluated + * @tparam T the clazz type + * + * @return the [[CaseClassDescriptor]] for the specified class + * + * @see [[CaseClassDescriptor]] + */ + def getConstraintsForClass[T](clazz: Class[T]): CaseClassDescriptor[T] = + descriptorFactory.describe[T](clazz) + + /** + * Returns a [[ExecutableDescriptor]] object describing a given [[Executable]]. + * + * The returned object (and associated objects including ExecutableDescriptors) are immutable. + * + * @param executable the [[Executable]] to describe. + * + * @return the [[ExecutableDescriptor]] for the specified [[Executable]] + * + * @note the returned [[ExecutableDescriptor]] is NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + private[twitter] def describeExecutable( + executable: Executable, + mixinClazz: Option[Class[_]] + ): ExecutableDescriptor = descriptorFactory.describeExecutable(executable, mixinClazz) + + /** + * Returns an array of [[MethodDescriptor]] instances describing the `@MethodValidation`-annotated + * or otherwise constrained methods of the given case class type. + * + * @param clazz class or interface type evaluated + * + * @return an array of [[MethodDescriptor]] for `@MethodValidation`-annotated or otherwise + * constrained methods of the class + * + * @note the returned [[MethodDescriptor]] instances are NOT cached. It is up to the caller of this + * method to optimize any calls to this method. + */ + private[twitter] def describeMethods( + clazz: Class[_] + ): Array[MethodDescriptor] = descriptorFactory.describeMethods(clazz) + + /** + * Returns the set of Constraints which validate the given [[Annotation]] class. + * + * For instance if the constraint type is `Class[com.twitter.util.validation.constraints.CountryCode]`, + * the returned set would contain an instance of the `ISO3166CountryCodeConstraintValidator`. + * + * @return the set of supporting constraint validators for a given [[Annotation]]. + * @note this method is memoized as it should only ever need to be calculated once for a given [[Class]]. + */ + val findConstraintValidators: Class[_ <: Annotation] => Set[ConstraintValidatorType] = Memoize { + annotationType: Class[_ <: Annotation] => + // compute from constraint, otherwise compute from registry + val validatedBy: Set[ConstraintValidatorType] = + if (annotationType.isAnnotationPresent(classOf[Constraint])) { + val constraintAnnotation = annotationType.getAnnotation(classOf[Constraint]) + constraintAnnotation + .validatedBy().map { clazz => + val instance = validatorFactory.constraintValidatorFactory.getInstance(clazz) + if (instance == null) + throw new ValidationException( + s"Constraint factory returned null when trying to create instance of $clazz.") + instance.asInstanceOf[ConstraintValidatorType] + }.toSet + } else Set.empty + if (validatedBy.isEmpty) { + validatorFactory.constraintHelper + .getAllValidatorDescriptors(annotationType).asScala + .map(_.newInstance(validatorFactory.constraintValidatorFactory)).toSet + } else validatedBy + } + + /** + * Checks whether the specified [[Annotation]] is a valid constraint constraint. A constraint + * constraint has to fulfill the following conditions: + * + * - Must be annotated with [[jakarta.validation.Constraint]] + * - Define a `message` parameter + * - Define a `groups` parameter + * - Define a `payload` parameter + * + * @param annotation The [[Annotation]] to test. + * + * @return true if the constraint fulfills the above conditions, false otherwise. + */ + def isConstraintAnnotation(annotation: Annotation): Boolean = + isConstraintAnnotation(annotation.annotationType()) + + /** + * Checks whether the specified [[Annotation]] clazz is a valid constraint constraint. A constraint + * constraint has to fulfill the following conditions: + * + * - Must be annotated with [[jakarta.validation.Constraint]] + * - Define a `message` parameter + * - Define a `groups` parameter + * - Define a `payload` parameter + * + * @return true if the constraint fulfills the above conditions, false otherwise. + * @note this method is memoized as it should only ever need to be calculated once for a given [[Class]]. + */ + val isConstraintAnnotation: Class[_ <: Annotation] => Boolean = Memoize { + annotationType: Class[_ <: Annotation] => + validatorFactory.constraintHelper.isConstraintAnnotation(annotationType) + } + + /** + * Validates all constraint constraints on an object. + * + * @param obj the object to validate. + * + * @throws ValidationException - if any constraints produce a violation. + * @throws IllegalArgumentException - if object is null. + */ + @throws[ValidationException] + def verify(obj: Any, groups: Class[_]*): Unit = { + val invalidResult = validate(obj, groups: _*) + if (invalidResult.nonEmpty) + throw constraintViolationFactory.newConstraintViolationException(invalidResult) + } + + /** + * Validates all constraint constraints on an object. + * + * @param obj the object to validate. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type of the object to validate. + * + * @return constraint violations or an empty set if none + * + * @throws IllegalArgumentException - if object is null. + */ + def validate[T]( + obj: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (obj == null) throw new IllegalArgumentException("value must not be null.") + val context = ValidationContext[T](None, None, None, None, PathImpl.createRootPath()) + validate[T]( + maybeDescriptor = None, + context = context, + value = obj, + groups = groups + ) + } + + /** + * Validates all constraint constraints on the `fieldName` field of the given class `beanType` + * if the `fieldName` field value were `value`. + * + * ConstraintViolation objects return `null` for ConstraintViolation.getRootBean() and + * ConstraintViolation.getLeafBean(). + * + * ==Usage== + * Validate value of "" for the "id" field in case class `MyClass` (enabling group "Checks"): + * {{{ + * case class MyCaseClass(@NotEmpty(groups = classOf[Checks]) id) + * validator.validateFieldValue(classOf[MyCaseClass], "id", "", Seq(classOf[Checks])) + * }}} + * + * @param beanType the case class type. + * @param propertyName field to validate. + * @param value field value to validate. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type of the object to validate + * + * @return constraint violations or an empty set if none + * + * @throws IllegalArgumentException - if beanType is null, if fieldName is null, empty or not a valid object property. + */ + def validateValue[T]( + beanType: Class[T], + propertyName: String, + value: Any, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (beanType == null) throw new IllegalArgumentException("beanType must not be null.") + val descriptor = getConstraintsForClass(beanType) + validateValue[T](descriptor, propertyName, value, groups: _*) + } + + /** + * Validates all constraint constraints on the `fieldName` field of the class described + * by the given [[CaseClassDescriptor]] if the `fieldName` field value were `value`. + * + * @param descriptor the [[CaseClassDescriptor]] of the described class. + * @param propertyName field to validate. + * @param value field value to validate. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type of the object to validate + * + * @return constraint violations or an empty set if none + * + * @throws IllegalArgumentException - if fieldName is null, empty or not a valid object property. + */ + def validateValue[T]( + descriptor: CaseClassDescriptor[T], + propertyName: String, + value: Any, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (descriptor == null) + throw new IllegalArgumentException("descriptor must not be null.") + if (propertyName == null || propertyName.isEmpty) + throw new IllegalArgumentException("fieldName must not be null or empty.") + + val mangledName = DescriptorFactory.mangleName(propertyName) + descriptor.members.get(mangledName) match { + case Some(propertyDescriptor) => + val context = ValidationContext( + Some(propertyName), + Some(descriptor.clazz), + None, + None, + PathImpl.createRootPath()) + validate( + maybeDescriptor = Some(propertyDescriptor), + context = context, + value = value, + groups = groups + ) + case _ => + throw new IllegalArgumentException(s"$propertyName is not a field of ${descriptor.clazz}.") + } + } + + /** + * Validates all constraint constraints on the `fieldName` field of the given object. + * + * ==Usage== + * Validate the "id" field in case class `MyClass` (enabling the validation group "Checks"): + * {{{ + * case class MyClass(@NotEmpty(groups = classOf[Checks]) id) + * val i = MyClass("") + * validator.validateProperty(i, "id", Seq(classOf[Checks])) + * }}} + * + * @param obj object to validate. + * @param propertyName property to validate (used in error reporting). + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type of the object to validate. + * + * @return constraint violations or an empty set if none. + * + * @throws IllegalArgumentException - if object is null, if fieldName is null, empty or not a valid object property. + */ + def validateProperty[T]( + obj: T, + propertyName: String, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (obj == null) throw new IllegalArgumentException("obj must not be null.") + if (propertyName == null || propertyName.isEmpty) + throw new IllegalArgumentException("fieldName must not be null or empty.") + + val caseClassDescriptor = getConstraintsForClass(obj.getClass) + + val mangledName = DescriptorFactory.mangleName(propertyName) + caseClassDescriptor.members.get(mangledName) match { + case Some(_) => + val context = ValidationContext( + Some(propertyName), + Some(obj.getClass.asInstanceOf[Class[T]]), + Some(obj), + Some(obj), + PathImpl.createRootPath()) + validate( + maybeDescriptor = Some(caseClassDescriptor), + context = context, + value = obj, + groups = groups + ) + case _ => + throw new IllegalArgumentException( + s"$propertyName is not a field of ${caseClassDescriptor.clazz}.") + } + } + + /** Returns the contract for validating parameters and return values of methods and constructors. */ + def forExecutables: ScalaExecutableValidator = this + + /** + * Returns an instance of the specified type allowing access to provider-specific APIs. + * + * If the Jakarta Bean Validation provider implementation does not support the specified class, ValidationException is thrown. + * + * @param clazz the class of the object to be returned + * @tparam U the type of the object to be returned + * + * @return an instance of the specified class + * + * @throws ValidationException - if the provider does not support the call. + */ + @throws[ValidationException] + def unwrap[U](clazz: Class[U]): U = { + // If this were an implementation that implemented the specification + // interface, this is where users would be able to cast the validator from the + // interface type into our specific implementation type in order to call methods + // only available on the implementation. However, since the specification uses Java + // collection types we do not directly implement and instead expose our feature-compatible + // ScalaValidator directly. We try to remain true to the spirit of the specification + // interface and thus implement unwrap but only to return this implementation. + if (clazz.isAssignableFrom(classOf[ScalaValidator])) { + this.asInstanceOf[U] + } else if (clazz.isAssignableFrom(classOf[ScalaExecutableValidator])) { + this.asInstanceOf[U] + } else { + throw new ValidationException(s"Type ${clazz.getName} not supported for unwrapping.") + } + } + + /* Executable Validator Methods */ + + /** @inheritdoc */ + def validateMethods[T]( + obj: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (obj == null) + throw new IllegalArgumentException("The object to be validated must not be null.") + val caseClassDescriptor = getConstraintsForClass(obj.getClass) + validateMethods[T]( + rootClazz = Some(obj.getClass.asInstanceOf[Class[T]]), + root = Some(obj), + methods = caseClassDescriptor.methods, + obj = obj, + groups = groups + ) + } + + /** @inheritdoc */ + def validateMethod[T]( + obj: T, + method: Method, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (obj == null) + throw new IllegalArgumentException("The object to be validated must not be null.") + if (method == null) + throw new IllegalArgumentException("The method to be validated must not be null.") + val caseClassDescriptor = getConstraintsForClass(obj.getClass) + caseClassDescriptor.methods.find(_.method.getName == method.getName) match { + case Some(methodDescriptor) => + validateMethods[T]( + rootClazz = Some(obj.getClass.asInstanceOf[Class[T]]), + root = Some(obj), + methods = Array(methodDescriptor), + obj = obj, + groups = groups + ) + case _ => + throw new IllegalArgumentException( + s"${method.getName} is not method of ${caseClassDescriptor.clazz}.") + } + } + + /** @inheritdoc */ + def validateParameters[T]( + obj: T, + method: Method, + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (obj == null) + throw new IllegalArgumentException("The object to be validated must not be null.") + if (method == null) + throw new IllegalArgumentException("The method to be validated must not be null.") + + val caseClassDescriptor = getConstraintsForClass(obj.getClass) + caseClassDescriptor.methods.find(_.method.getName == method.getName) match { + case Some(methodDescriptor: ExecutableDescriptor) => + validateExecutableParameters[T]( + executableDescriptor = methodDescriptor, + root = Some(obj), + leaf = Some(obj), + parameterValues = parameterValues, + parameterNames = None, + groups = groups + ) + case _ => + throw new IllegalArgumentException( + s"${method.getName} is not method of ${caseClassDescriptor.clazz}.") + } + } + + /** @inheritdoc */ + def validateReturnValue[T]( + obj: T, + method: Method, + returnValue: Any, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = underlying + .forExecutables().validateReturnValue( + obj, + method, + returnValue, + groups: _* + ).asScala.toSet + + /** @inheritdoc */ + def validateConstructorParameters[T]( + constructor: Constructor[T], + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (constructor == null) + throw new IllegalArgumentException("The executable to be validated must not be null.") + if (parameterValues == null) + throw new IllegalArgumentException("The constructor parameter array cannot not be null.") + + val parameterCount = constructor.getParameterCount + val parameterValuesLength = parameterValues.length + if (parameterCount != parameterValuesLength) + throw new ValidationException( + s"Wrong number of parameters. Method or constructor $constructor expects $parameterCount parameters, but got $parameterValuesLength.") + + validateExecutableParameters[T]( + executableDescriptor = descriptorFactory.describe[T](constructor), + root = None, + leaf = None, + parameterValues = parameterValues, + parameterNames = None, + groups = groups + ) + } + + /** @inheritdoc */ + def validateMethodParameters[T]( + method: Method, + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (method == null) + throw new IllegalArgumentException("The executable to be validated must not be null.") + if (parameterValues == null) + throw new IllegalArgumentException("The method parameter array cannot not be null.") + + val parameterCount = method.getParameterCount + val parameterValuesLength = parameterValues.length + if (parameterCount != parameterValuesLength) + throw new ValidationException( + s"Wrong number of parameters. Method or constructor $method expects $parameterCount parameters, but got $parameterValuesLength.") + + descriptorFactory.describe(method) match { + case Some(descriptor: ExecutableDescriptor) => + validateExecutableParameters[T]( + executableDescriptor = descriptor, + root = None, + leaf = None, + parameterValues = parameterValues, + parameterNames = None, + groups = groups + ) + case _ => Set.empty[ConstraintViolation[T]] + } + } + + /** @inheritdoc */ + def validateConstructorReturnValue[T]( + constructor: Constructor[T], + createdObject: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = + underlying + .forExecutables().validateConstructorReturnValue( + constructor, + createdObject, + groups: _*).asScala.toSet + + /* Private */ + + /** @inheritdoc */ + private[twitter] def validateExecutableParameters[T]( + executable: ExecutableDescriptor, + parameterValues: Array[Any], + parameterNames: Array[String], + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + if (executable == null) + throw new IllegalArgumentException("The executable to be validated must not be null.") + if (parameterValues == null) + throw new IllegalArgumentException("The constructor parameter array cannot not be null.") + + val parameterCount = executable.executable.getParameterCount + val parameterValuesLength = parameterValues.length + if (parameterCount != parameterValuesLength) + throw new ValidationException( + s"Wrong number of parameters. Method or constructor $executable expects $parameterCount parameters, but got $parameterValuesLength.") + + validateExecutableParameters[T]( + executableDescriptor = executable, + root = None, + leaf = None, + parameterValues = parameterValues, + parameterNames = Some(parameterNames), + groups = groups + ) + } + + /** @inheritdoc */ + private[twitter] def validateMethods[T]( + methods: Array[MethodDescriptor], + obj: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] = { + validateMethods[T]( + rootClazz = Some(obj.getClass.asInstanceOf[Class[T]]), + root = Some(obj), + methods = methods, + obj = obj, + groups = groups + ) + } + + // note: Prefer while loop over Scala for loop for better performance. The scala for loop + // performance is optimized in 2.13.0 if we enable scalac: https://github.com/scala/bug/issues/1338 + + /** + * Evaluates all known [[ConstraintValidator]] instances supporting the constraints represented by the + * given mapping of [[Annotation]] to constraint attributes placed on a field `fieldName` if + * the `fieldName` field value were `value`. The `fieldName` is used in error reporting (as the + * the resulting property path for any resulting [[ConstraintViolation]]). + * + * E.g., the given mapping could be Map(jakarta.validation.constraint.Size -> Map("min" -> 0, "max" -> 1000)) + * + * ConstraintViolation objects return `null` for ConstraintViolation.getRootBeanClass(), + * ConstraintViolation.getRootBean() and ConstraintViolation.getLeafBean(). + * + * @param constraints the mapping of constraint [[Annotation]] to attributes to evaluate + * @param fieldName field to validate (used in error reporting). + * @param value field value to validate. + * @param groups the list of groups targeted for validation (defaults to Default). + * + * @return constraint violations or an empty set if none + * + * @throws IllegalArgumentException - if fieldName is null, empty or not a valid object property. + * + * @note Scala users should prefer the version which takes a [[Map]] of constraints and + * a [[Seq]] of groups. + */ + @varargs + private[twitter] def validateFieldValue( + constraints: java.util.Map[Class[_ <: Annotation], java.util.Map[String, Any]], + fieldName: String, + value: Any, + groups: Class[_]* + ): java.util.Set[ConstraintViolation[Any]] = { + if (fieldName == null || fieldName.isEmpty) + throw new IllegalArgumentException("fieldName must not be null or empty.") + val annotations = constraints.asScala.map { + case (constraintAnnotationType, attributes) + if isConstraintAnnotation(constraintAnnotationType) => + AnnotationFactory.newInstance(constraintAnnotationType, attributes.asScala.toMap) + } + validateFieldValue(fieldName, annotations.toArray, value, groups.toSeq).asJava + } + + /** + * Evaluates all known [[ConstraintValidator]] instances supporting the constraints represented by the + * given mapping of [[Annotation]] to constraint attributes placed on a field `fieldName` if + * the `fieldName` field value were `value`. The `fieldName` is used in error reporting (as the + * the resulting property path for any resulting [[ConstraintViolation]]). + * + * E.g., the given mapping could be Map(jakarta.validation.constraint.Size -> Map("min" -> 0, "max" -> 1000)) + * + * ConstraintViolation objects return `null` for ConstraintViolation.getRootBeanClass(), + * ConstraintViolation.getRootBean() and ConstraintViolation.getLeafBean(). + * + * @param constraints the mapping of constraint [[Annotation]] to attributes to evaluate + * @param fieldName field to validate (used in error reporting). + * @param value field value to validate. + * @param groups the list of groups targeted for validation (defaults to Default). + * + * @return constraint violations or an empty set if none + * + * @throws IllegalArgumentException - if fieldName is null, empty or not a valid object property. + * + * @note Java users should prefer the version which takes a [[java.util.Map]] of constraints and + * a [[java.util.List]] of groups. + */ + private[twitter] def validateFieldValue( + constraints: Map[Class[_ <: Annotation], Map[String, Any]], + fieldName: String, + value: Any, + groups: Class[_]* + ): Set[ConstraintViolation[Any]] = { + if (fieldName == null || fieldName.isEmpty) + throw new IllegalArgumentException("fieldName must not be null or empty.") + + val annotations = constraints.map { + case (constraintAnnotationType, attributes) => + AnnotationFactory.newInstance(constraintAnnotationType, attributes) + } + + validateFieldValue(fieldName, annotations.toArray, value, groups) + } + + /** Validate the given fieldName with the given constraint constraints, value, and groups. */ + private[this] def validateFieldValue( + fieldName: String, + constraints: Array[Annotation], + value: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[Any]] = + if (constraints.nonEmpty) { + val results = new mutable.ListBuffer[ConstraintViolation[Any]]() + var index = 0 + val length = constraints.length + while (index < length) { + val annotation = constraints(index) + val path = PathImpl.createRootPath + path.addPropertyNode(fieldName) + val clazz = classOf[Any] + val scalaType = Reflector.scalaTypeOf(value.getClass) + val context = ValidationContext(Some(fieldName), Some(clazz), None, None, path) + + results.appendAll( + isValid[Any]( + context = context, + constraint = annotation, + scalaType = scalaType, + value = value, + groups = groups + ) + ) + index += 1 + } + results.toSet + } else Set.empty[ConstraintViolation[Any]] + + // https://beanvalidation.org/2.0/spec/#constraintsdefinitionimplementation-validationimplementation + // From the documentation: + // + // While not mandatory, it is considered a good practice to split the core constraint + // validation from the "not null" constraint validation (for example, an @Email constraint + // should return "true" on a `null` object, i.e. it will not also assert the @NotNull validation). + // A "null" can have multiple meanings but is commonly used to express that a value does + // not make sense, is not available or is simply unknown. Those constraints on the + // value are orthogonal in most cases to other constraints. For example a String, + // if present, must be an email but can be null. Separating both concerns is a good + // practice. + // + // Thus, we "ignore" any value that is null and not annotated with @NotNull. + private[this] def ignorable(value: Any, annotation: Annotation): Boolean = value match { + case null if !ReflectAnnotations.equals[jakarta.validation.constraints.NotNull](annotation) => + true + case _ => false + } + + /** Validate Executable field-level parameters (including cross-field parameters) */ + private[this] def validateExecutableParameters[T]( + executableDescriptor: ExecutableDescriptor, + root: Option[T], + leaf: Option[Any], + parameterValues: Array[Any], + parameterNames: Option[Array[String]], + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = { + if (executableDescriptor.members.size != parameterValues.length) + throw new IllegalArgumentException( + s"Invalid number of arguments for method ${executableDescriptor.executable.getName}.") + val results = new mutable.ListBuffer[ConstraintViolation[T]]() + val parameters: Array[Parameter] = executableDescriptor.executable.getParameters + val parameterNamesList: Array[String] = + (parameterNames match { + case Some(names) => + names + case _ => + getExecutableParameterNames(executableDescriptor.executable) + }).map { maybeMangled => + DescriptorFactory.unmangleName( + maybeMangled + ) // parameter names are encoded since Scala 2.13.5 + } + + // executable parameter constraints + var index = 0 + val parametersLength = parameterNamesList.length + while (index < parametersLength) { + val parameter = parameters(index) + val parameterValue = parameterValues(index) + // only for property path use to affect error reporting + val parameterName = parameterNamesList(index) + + val parameterPath = + PathImpl.createPathForExecutable(getExecutableMetaData(executableDescriptor.executable)) + parameterPath.addParameterNode(parameterName, index) + + val propertyDescriptor = executableDescriptor.members(parameter.getName) + val validateFieldContext = + ValidationContext[T]( + Some(parameterName), + Some(executableDescriptor.executable.getDeclaringClass.asInstanceOf[Class[T]]), + root, + leaf, + parameterPath) + val parameterViolations = validateField[T]( + validateFieldContext, + propertyDescriptor, + parameterValue, + groups + ) + if (parameterViolations.nonEmpty) results.appendAll(parameterViolations) + + index += 1 + } + results.toSet + + // executable cross-parameter constraints + val crossParameterPath = PathImpl + .createPathForExecutable(getExecutableMetaData(executableDescriptor.executable)) + crossParameterPath.addCrossParameterNode() + val constraints = executableDescriptor.annotations + index = 0 + val constraintsLength = constraints.length + while (index < constraintsLength) { + val constraint = constraints(index) + val scalaType = Reflector.scalaTypeOf(parameterValues.getClass) + + val validators = + findConstraintValidators(constraint.annotationType()) + .filter(parametersValidationTargetFilter) + + val validatorsIterator = validators.iterator + while (validatorsIterator.hasNext) { + val validator = validatorsIterator.next().asInstanceOf[ConstraintValidator[Annotation, _]] + val constraintDescriptor = constraintDescriptorFactory + .newConstraintDescriptor( + name = executableDescriptor.executable.getName, + clazz = scalaType.erasure, + declaringClazz = executableDescriptor.executable.getDeclaringClass, + annotation = constraint, + constrainedElementKind = ConstrainedElementKind.METHOD + ) + if (groupsEnabled(constraintDescriptor, groups)) { + // initialize validator + validator.initialize(constraint) + val validateCrossParameterContext = + ValidationContext[T]( + None, + Some(executableDescriptor.executable.getDeclaringClass.asInstanceOf[Class[T]]), + root, + leaf, + crossParameterPath) + results.appendAll( + isValid( + context = validateCrossParameterContext, + constraint = constraint, + scalaType = scalaType, + value = parameterValues, + constraintDescriptorOption = Some(constraintDescriptor), + validatorOption = Some(validator.asInstanceOf[ConstraintValidator[Annotation, Any]]), + groups = groups + ) + ) + } + } + index += 1 + } + results.toSet + } + + /** Validate method-level constraints, i.e., @MethodValidation annotated methods */ + private[this] def validateMethods[T]( + rootClazz: Option[Class[T]], + root: Option[T], + methods: Array[MethodDescriptor], + obj: T, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = { + val results = new mutable.ListBuffer[ConstraintViolation[T]]() + if (methods.nonEmpty) { + var index = 0 + val length = methods.length + while (index < length) { + val methodDescriptor = methods(index) + val validateMethodContext = + ValidationContext(None, rootClazz, root, Some(obj), PathImpl.createRootPath()) + results.appendAll( + validateMethod(validateMethodContext, methodDescriptor, obj, groups) + ) + index += 1 + } + } + results.toSet + } + + // BEGIN: Recursive validation methods ----------------------------------------------------------- + + /** Validate field-level constraints */ + private[this] def validateField[T]( + context: ValidationContext[T], + propertyDescriptor: PropertyDescriptor, + fieldValue: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = { + val constraints: Array[Annotation] = propertyDescriptor.annotations + val results = new mutable.ListBuffer[ConstraintViolation[T]]() + var index = 0 + val length = constraints.length + while (index < length) { + val constraint = constraints(index) + results.appendAll( + isValid[T]( + context, + constraint = constraint, + scalaType = propertyDescriptor.scalaType, + value = fieldValue, + groups = groups + ) + ) + index += 1 + } + + // Cannot cascade a null or None value + fieldValue match { + case null | None => // do nothing + case _ => + if (propertyDescriptor.isCascaded) { + results.appendAll( + validateCascadedProperty[T](context, propertyDescriptor, fieldValue, groups) + ) + } + } + + if (results.isEmpty) Set.empty[ConstraintViolation[T]] + else results.toSet + } + + /** Validate cascaded field-level properties */ + private[this] def validateCascadedProperty[T]( + context: ValidationContext[T], + propertyDescriptor: PropertyDescriptor, + clazzInstance: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = propertyDescriptor.cascadedScalaType match { + case Some(cascadedScalaType) => + // Update the found value with the current landscape, for example, the retrieved case class + // descriptor may be scala type Foo, but this property is a Seq[Foo] or an Option[Foo] and + // we want to carry that property typing through. + val caseClassDescriptor = + descriptorFactory + .describe[T](cascadedScalaType.erasure.asInstanceOf[Class[T]]) + .copy(scalaType = propertyDescriptor.scalaType) + + val caseClassPath = PathImpl.createCopy(context.path) + caseClassDescriptor.scalaType match { + case argType if argType.isCollection => + val results = mutable.ListBuffer[ConstraintViolation[T]]() + val collectionValue: Iterable[_] = clazzInstance.asInstanceOf[Iterable[_]] + val collectionValueIterator = collectionValue.iterator + var index = 0 + while (collectionValueIterator.hasNext) { + val instanceValue = collectionValueIterator.next() + // apply the index to the parent path, then use this to recompute paths of members and methods + val indexedPropertyPath = { + val indexedPath = PathImpl.createCopyWithoutLeafNode(caseClassPath) + indexedPath.addPropertyNode( + s"${caseClassPath.getLeafNode.asString()}[${index.toString}]") + indexedPath + } + val violations = validate[T]( + maybeDescriptor = Some( + caseClassDescriptor + .copy( + scalaType = caseClassDescriptor.scalaType.typeArgs.head, + members = caseClassDescriptor.members, + methods = caseClassDescriptor.methods + ) + ), + context = context.copy(path = indexedPropertyPath), + value = instanceValue, + groups = groups + ) + if (violations.nonEmpty) results.appendAll(violations) + index += 1 + } + if (results.nonEmpty) results.toSet + else Set.empty[ConstraintViolation[T]] + case argType if argType.isOption => + val optionValue = clazzInstance.asInstanceOf[Option[_]] + validate[T]( + maybeDescriptor = Some( + caseClassDescriptor.copy(scalaType = caseClassDescriptor.scalaType.typeArgs.head)), + context = context.copy(path = caseClassPath), + value = optionValue, + groups = groups + ) + case _ => + validate[T]( + maybeDescriptor = Some(caseClassDescriptor), + context = context.copy(path = caseClassPath), + value = clazzInstance, + groups = groups + ) + } + case _ => Set.empty[ConstraintViolation[T]] + } + + /** Validate @MethodValidation annotated methods */ + private[this] def validateMethod[T]( + context: ValidationContext[T], + methodDescriptor: MethodDescriptor, + clazzInstance: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = + ReflectAnnotations.findAnnotation[MethodValidation](methodDescriptor.annotations) match { + case Some(annotation) => + // if we're unable to invoke the method, we want this propagate + val constraintDescriptor = new MethodValidationConstraintDescriptor(annotation) + if (groupsEnabled(constraintDescriptor, groups)) { + try { + val result = + methodDescriptor.method.invoke(clazzInstance).asInstanceOf[MethodValidationResult] + val pathWithMethodName = PathImpl.createCopy(context.path) + pathWithMethodName.addPropertyNode(methodDescriptor.method.getName) + parseMethodValidationFailures[T]( + rootClazz = context.rootClazz, + root = context.root, + leaf = Some(clazzInstance), + path = pathWithMethodName, + annotation = annotation.asInstanceOf[MethodValidation], + method = methodDescriptor.method, + result = result, + value = clazzInstance, + constraintDescriptor + ) + } catch { + case e: InvocationTargetException if e.getCause != null => + throw e.getCause + } + } else Set.empty[ConstraintViolation[T]] + case _ => Set.empty[ConstraintViolation[T]] + } + + /** Validate class-level constraints */ + private[this] def validateClazz[T]( + context: ValidationContext[T], + caseClassDescriptor: CaseClassDescriptor[T], + clazzInstance: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = { + val constraints = caseClassDescriptor.annotations + if (constraints.nonEmpty) { + val results = new mutable.ListBuffer[ConstraintViolation[T]]() + var index = 0 + val length = constraints.length + while (index < length) { + val annotation = constraints(index) + results.appendAll( + isValid[T]( + context = context, + constraint = annotation, + scalaType = caseClassDescriptor.scalaType, + value = clazzInstance, + groups = groups + ) + ) + index += 1 + } + results.toSet + } else Set.empty[ConstraintViolation[T]] + } + + /** Recursively validate full case class */ + @tailrec + private[this] def validate[T]( + maybeDescriptor: Option[Descriptor], + context: ValidationContext[T], + value: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = maybeDescriptor match { + case Some(descriptor) => + descriptor match { + case p: PropertyDescriptor => + val propertyPath = PathImpl.createCopy(context.path) + context.fieldName.map(propertyPath.addPropertyNode) + // validateField recurses back through validate for cascaded properties + validateField[T](context.copy(path = propertyPath), p, value, groups) + case c: CaseClassDescriptor[_] => + val memberViolations = { + val results = new mutable.ListBuffer[ConstraintViolation[T]]() + if (c.members.nonEmpty) { + val keys: Iterable[String] = c.members.keys + val keyIterator = keys.iterator + while (keyIterator.hasNext) { + val name = keyIterator.next() + val propertyDescriptor = c.members(name) + val propertyPath = PathImpl.createCopy(context.path) + propertyPath.addPropertyNode(DescriptorFactory.unmangleName(name)) + // validateField recurses back through validate here for cascaded properties + val fieldResults = + validateField[T]( + context.copy(fieldName = Some(name), path = propertyPath), + propertyDescriptor, + findFieldValue(value, c.clazz, name), + groups) + if (fieldResults.nonEmpty) results.appendAll(fieldResults) + } + } + results.toSet + } + val methodViolations = { + val results = new ListBuffer[ConstraintViolation[T]]() + if (c.methods.nonEmpty) { + var index = 0 + val length = c.methods.length + while (index < length) { + val methodResults = value match { + case null | None => + Set.empty + case Some(obj) => + validateMethod[T](context.copy(fieldName = None), c.methods(index), obj, groups) + case _ => + validateMethod[T]( + context.copy(fieldName = None), + c.methods(index), + value, + groups) + } + + if (methodResults.nonEmpty) results.appendAll(methodResults) + index += 1 + } + } + results.toSet + } + val clazzViolations = + validateClazz[T](context, c.asInstanceOf[CaseClassDescriptor[T]], value, groups) + + // put them all together + memberViolations ++ methodViolations ++ clazzViolations + } + case _ => + val clazz: Class[T] = value.getClass.asInstanceOf[Class[T]] + if (ReflectTypes.notCaseClass(clazz)) + throw new ValidationException(s"$clazz is not a valid case class.") + val caseClassDescriptor: CaseClassDescriptor[T] = descriptorFactory.describe(clazz) + val context = ValidationContext( + fieldName = None, + rootClazz = Some(clazz), + root = Some(value.asInstanceOf[T]), + leaf = Some(value), + path = PathImpl.createRootPath()) + validate[T](Some(caseClassDescriptor), context, value, groups) + } + + // END: Recursive validation methods ------------------------------------------------------------- + + @tailrec + private[this] def findFieldValue[T]( + obj: Any, + clazz: Class[T], + name: String + ): Any = obj match { + case null | None => + obj + case Some(v) => + findFieldValue(v, clazz, name) + case _ => + try { + val field = clazz.getDeclaredField(name) + field.setAccessible(true) + Option(field.get(obj)).orNull.asInstanceOf[T] + } catch { + case NonFatal(t) => + throw new ValidationException(t) + } + } + + private[this] def isValid[T]( + context: ValidationContext[T], + constraint: Annotation, + scalaType: ScalaType, + value: Any, + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = + isValid(context, constraint, scalaType, value, None, None, groups) + + private[this] def isValid[T]( + context: ValidationContext[T], + constraint: Annotation, + scalaType: ScalaType, + value: Any, + constraintDescriptorOption: Option[ConstraintDescriptorImpl[Annotation]], + validatorOption: Option[ConstraintValidator[Annotation, Any]], + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = value match { + case _: Option[_] => + isValidOption( + context, + constraint, + scalaType, + value, + constraintDescriptorOption, + validatorOption, + groups) + case _ => + val constraintDescriptor = constraintDescriptorOption match { + case Some(descriptor) => + descriptor + case _ => + constraintDescriptorFactory.newConstraintDescriptor( + name = context.fieldName.orNull, + clazz = scalaType.erasure, + declaringClazz = context.rootClazz.getOrElse(scalaType.erasure), + annotation = constraint + ) + } + + if (!ignorable(value, constraint) && groupsEnabled(constraintDescriptor, groups)) { + val refinedScalaType = + Types.refineScalaType(value, scalaType) + + val constraintValidator: ConstraintValidator[Annotation, Any] = validatorOption match { + case Some(validator) => + validator // should already be initialized + case _ => + constraintValidatorManager + .getInitializedValidator[Annotation]( + Types.getJavaType(refinedScalaType), + constraintDescriptor, + validatorFactory.constraintValidatorFactory, + validatorFactory.validatorFactoryScopedContext.getConstraintValidatorInitializationContext + ).asInstanceOf[ConstraintValidator[Annotation, Any]] + } + + if (constraintValidator == null) { + val configuration = + if (context.path.toString.isEmpty) Classes.simpleName(scalaType.erasure) + else context.path.toString + throw new UnexpectedTypeException( + s"No validator could be found for constraint '${constraint.annotationType()}' " + + s"validating type '${scalaType.erasure.getName}'. " + + s"Check configuration for '$configuration'") + } + // create validator context + val constraintValidatorContext: ConstraintValidatorContext = + constraintValidatorContextFactory.newConstraintValidatorContext( + context.path, + constraintDescriptor) + // compute if valid + if (constraintValidator.isValid(value, constraintValidatorContext)) Set.empty + else { + constraintViolationFactory.buildConstraintViolations[T]( + rootClazz = context.rootClazz, + root = context.root, + leaf = context.leaf, + path = context.path, + invalidValue = value, + constraintDescriptor = constraintDescriptor, + constraintValidatorContext = constraintValidatorContext + ) + } + } else Set.empty + } + + private[this] def isValidOption[T]( + context: ValidationContext[T], + constraint: Annotation, + scalaType: ScalaType, + value: Any, + constraintDescriptorOption: Option[ConstraintDescriptorImpl[Annotation]], + validatorOption: Option[ConstraintValidator[Annotation, Any]], + groups: Seq[Class[_]] + ): Set[ConstraintViolation[T]] = value match { + case Some(actualVal) => + isValid[T]( + context, + constraint, + scalaType, + actualVal, + constraintDescriptorOption, + validatorOption, + groups) + case _ => + Set.empty[ConstraintViolation[T]] + } + + private[this] def groupsEnabled( + constraintDescriptor: ConstraintDescriptor[_], + groups: Seq[Class[_]] + ): Boolean = { + val groupsFromAnnotation: java.util.Set[Class[_]] = constraintDescriptor.getGroups + if (groups.isEmpty && groupsFromAnnotation.isEmpty) true + else if (groups.isEmpty && + groupsFromAnnotation.contains(classOf[jakarta.validation.groups.Default])) true + else if (groups.contains(classOf[jakarta.validation.groups.Default]) && + groupsFromAnnotation.isEmpty) true + else groups.exists(groupsFromAnnotation.contains) + } + + /** @note this method is memoized as it should only ever need to be calculated once for a given [[Executable]] */ + private[this] val getExecutableParameterNames: Executable => Array[String] = Memoize { + executable: Executable => + val parameterNameProviderNames: java.util.List[String] = + validatorFactory.validatorFactoryScopedContext.getParameterNameProvider + .getParameterNames(executable) + // depending on the underlying list this may be linear and not constant lookup + Array.tabulate(parameterNameProviderNames.size()) { index => + parameterNameProviderNames.get(index) + } + } + + /** @note this method is memoized as it should only ever need to be calculated once for a given [[ConstraintValidator]] */ + private[this] val parametersValidationTargetFilter: ConstraintValidator[_, _] => Boolean = + Memoize { constraintValidator: ConstraintValidator[_, _] => + ReflectAnnotations + .findAnnotation( + classOf[SupportedValidationTarget], + constraintValidator.getClass.getAnnotations) match { + case Some(annotation: SupportedValidationTarget) => + annotation.value().contains(ValidationTarget.PARAMETERS) + case _ => false + } + } + + /** @note this method is memoized as it should only ever need to be calculated once for a given [[Executable]] */ + private[this] val getExecutableMetaData: Executable => ExecutableMetaData = Memoize { + executable: Executable => + val callable: ExecutableCallable[_] = executable match { + case constructor: Constructor[_] => + new ConstructorCallable(constructor) + case method: Method => + new MethodCallable(method) + } + val builder = new ExecutableMetaData.Builder( + executable.getDeclaringClass, + new ConstrainedExecutable( + ConfigurationSource.ANNOTATION, + callable, + callable.getParameters + .map(p => + new ConstrainedParameter( + ConfigurationSource.ANNOTATION, + callable, + p.getType, + p.getIndex)).asJava, + Collections.emptySet(), + Collections.emptySet(), + Collections.emptySet(), + CascadingMetaDataBuilder.nonCascading + ), + validatorFactory.getConstraintCreationContext, + validatorFactory.getExecutableHelper, + validatorFactory.getExecutableParameterNameProvider, + validatorFactory.getMethodValidationConfiguration + ) + + builder.build() + } + + private[validation] def parseMethodValidationFailures[T]( + rootClazz: Option[Class[T]], + root: Option[T], + leaf: Option[Any], + path: PathImpl, + annotation: MethodValidation, + method: Method, + result: MethodValidationResult, + value: Any, + constraintDescriptor: ConstraintDescriptor[_] + ): Set[ConstraintViolation[T]] = { + def methodValidationConstraintViolation( + validationPath: PathImpl, + message: String, + payload: Payload + ): ConstraintViolation[T] = { + constraintViolationFactory.newConstraintViolation[T]( + messageTemplate = annotation.message(), + interpolatedMessage = message, + path = validationPath, + invalidValue = value.asInstanceOf[T], + rootClazz = rootClazz.orNull, + root = root.getOrElse(null.asInstanceOf[T]), + leaf = leaf.orNull, + constraintDescriptor = constraintDescriptor, + payload = payload + ) + } + + result match { + case invalid: MethodValidationResult.Invalid => + val results = mutable.HashSet[ConstraintViolation[T]]() + val annotationFields: Array[String] = annotation.fields.filter(_.nonEmpty) + if (annotationFields.nonEmpty) { + var index = 0 + val length = annotationFields.length + while (index < length) { + val fieldName = annotationFields(index) + val parameterPath = PathImpl.createCopy(path) + parameterPath.addParameterNode(fieldName, index) + results.add( + methodValidationConstraintViolation( + parameterPath, + invalid.message, + invalid.payload.orNull) + ) + index += 1 + } + if (results.nonEmpty) results.toSet + else Set.empty[ConstraintViolation[T]] + } else { + Set(methodValidationConstraintViolation(path, invalid.message, invalid.payload.orNull)) + } + case _ => Set.empty[ConstraintViolation[T]] + } + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/cfg/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/cfg/BUILD new file mode 100644 index 0000000000..28fe3eb9f5 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/cfg/BUILD @@ -0,0 +1,19 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-cfg", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/cfg/ConstraintMapping.scala b/util-validator/src/main/scala/com/twitter/util/validation/cfg/ConstraintMapping.scala new file mode 100644 index 0000000000..ae8bb8eac6 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/cfg/ConstraintMapping.scala @@ -0,0 +1,41 @@ +package com.twitter.util.validation.cfg + +import jakarta.validation.ConstraintValidator +import java.lang.annotation.Annotation + +/** + * Simple configuration class for defining constraint mappings for ScalaValidator configuration. + * ==Usage== + * {{{ + * val customConstraintMapping: ConstraintMapping = + * ConstraintMapping(classOf[Annotation], classOf[ConstraintValidator]) + * + * val validator: ScalaValidator = + * ScalaValidator.builder + * .withConstraintMapping(customConstraintMapping) + * .validator + * }}} + * + * or multiple mappings + * + * {{{ + * val customConstraintMapping1: ConstraintMapping = ??? + * val customConstraintMapping2: ConstraintMapping = ??? + * + * val validator: ScalaValidator = + * ScalaValidator.builder + * .withConstraintMappings(Set(customConstraintMapping1, customConstraintMapping2)) + * .validator + * }}} + * + * @param annotationType the `Class[Annotation]` of the constraint annotation. + * @param constraintValidator the implementing ConstraintValidator class for the given constraint annotation. + * @param includeExistingValidators if this is an additional validator for the given constraint annotation type + * or if this should replace all existing validators for the given constraint annotation + * type. Default is true (additional). + * @note adding multiple constraint mappings for the same annotation type will result in a ValidationException being thrown. + */ +case class ConstraintMapping( + annotationType: Class[_ <: Annotation], + constraintValidator: Class[_ <: ConstraintValidator[_ <: Annotation, _]], + includeExistingValidators: Boolean = true) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/BUILD new file mode 100644 index 0000000000..4adad13620 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/BUILD @@ -0,0 +1,18 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-constraintvalidation", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-slf4j-api/src/main/scala", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/TwitterConstraintValidatorContext.scala b/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/TwitterConstraintValidatorContext.scala new file mode 100644 index 0000000000..4894abd34f --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation/TwitterConstraintValidatorContext.scala @@ -0,0 +1,221 @@ +package com.twitter.util.validation.constraintvalidation + +import jakarta.validation.{ConstraintValidatorContext, Payload, ValidationException} +import org.hibernate.validator.constraintvalidation.HibernateConstraintValidatorContext + +object TwitterConstraintValidatorContext { + + def builder: TwitterConstraintValidatorContext.Builder = + new Builder(None, None, Map.empty[String, AnyRef], Map.empty[String, AnyRef]) + + /** + * Simple utility to build a non-default [[jakarta.validation.ConstraintViolation]] + * from a [[jakarta.validation.ConstraintValidatorContext]]. + */ + class Builder private[validation] ( + messageTemplate: Option[String], + payload: Option[Payload], + messageParameters: Map[String, AnyRef], + expressionVariables: Map[String, AnyRef]) { + + /** + * Specify a new un-interpolated constraint message. + * + * @param messageTemplate an un-interpolated constraint message. + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def withMessageTemplate(messageTemplate: String): Builder = + new Builder( + messageTemplate = Some(messageTemplate), + payload = this.payload, + messageParameters = this.messageParameters, + expressionVariables = this.expressionVariables) + + /** + * Allows to set an object that may further describe the violation. The user is responsible to + * ensure that this payload is serializable in case the [[jakarta.validation.ConstraintViolation]] + * has to be serialized. + * + * @param payload payload an object representing additional information about the violation + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def withDynamicPayload(payload: Payload): Builder = + new Builder( + messageTemplate = this.messageTemplate, + payload = Some(payload), + messageParameters = this.messageParameters, + expressionVariables = this.expressionVariables) + + /** + * Allows setting of an additional named parameter which can be interpolated in the constraint violation message. The + * variable will be available for interpolation for all constraint violations generated for this constraint. + * This includes the default one as well as all violations created by this TwitterConstraintValidatorContext. + * To create multiple constraint violations with different variable values, this method can be called + * between successive calls to [[TwitterConstraintValidatorContext#addConstraintViolation]]. + * + * @param name the name under which to bind the parameter, cannot be `null`. + * @param value the value to be bound to the specified name + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def addMessageParameter(name: String, value: AnyRef): Builder = + new Builder( + messageTemplate = this.messageTemplate, + payload = this.payload, + messageParameters = this.messageParameters + (name -> value), + expressionVariables = this.expressionVariables + ) + + /** + * Allows setting of an additional expression variable which will be available as an EL variable during interpolation. The + * variable will be available for interpolation for all constraint violations generated for this constraint. + * This includes the default one as well as all violations created by this TwitterConstraintValidatorContext. + * To create multiple constraint violations with different variable values, this method can be called + * between successive calls to [[TwitterConstraintValidatorContext#addConstraintViolation]]. + * + * @param name the name under which to bind the expression variable, cannot be `null`. + * @param value the value to be bound to the specified name + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def addExpressionVariable(name: String, value: AnyRef): Builder = + new Builder( + messageTemplate = this.messageTemplate, + payload = this.payload, + messageParameters = this.messageParameters, + expressionVariables = this.expressionVariables + (name -> value) + ) + + /** + * Replaces the default constraint violation report with one built via this method using the + * given [[TwitterConstraintValidatorContext]]. + * + * @param context the [[TwitterConstraintValidatorContext]] to use in building a [[jakarta.validation.ConstraintViolation]] + * + * @throws ValidationException - if the given [[TwitterConstraintValidatorContext]] cannot be unwrapped as expected. + */ + @throws[ValidationException] + def addConstraintViolation(context: ConstraintValidatorContext): Unit = { + if (context != null) { + val hibernateContext: HibernateConstraintValidatorContext = + context.unwrap(classOf[HibernateConstraintValidatorContext]) + // custom payload + payload.foreach(hibernateContext.withDynamicPayload) + // set message parameters + messageParameters.foreach { + case (name, value) => + hibernateContext.addMessageParameter(name, value) + } + // set expression variables + expressionVariables.foreach { + case (name, value) => + hibernateContext.addExpressionVariable(name, value) + } + // custom message + messageTemplate.foreach { template => + hibernateContext.disableDefaultConstraintViolation() + + if (expressionVariables.nonEmpty) { + hibernateContext + .buildConstraintViolationWithTemplate(template) + .enableExpressionLanguage + .addConstraintViolation() + } else { + hibernateContext + .buildConstraintViolationWithTemplate(template) + .addConstraintViolation() + } + } + } + } + } + + /** + * Specify a new un-interpolated constraint message. + * + * @param messageTemplate an un-interpolated constraint message. + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def withMessageTemplate(messageTemplate: String): TwitterConstraintValidatorContext.Builder = + builder.withMessageTemplate(messageTemplate) + + /** + * Allows to set an object that may further describe the violation. The user is responsible to + * ensure that this payload is serializable in case the [[jakarta.validation.ConstraintViolation]] + * has to be serialized. + * + * @param payload payload an object representing additional information about the violation + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def withDynamicPayload(payload: Payload): TwitterConstraintValidatorContext.Builder = + builder.withDynamicPayload(payload) + + /** + * Allows setting of an additional named parameter which can be interpolated in the constraint violation message. The + * variable will be available for interpolation for all constraint violations generated for this constraint. + * This includes the default one as well as all violations created by this TwitterConstraintValidatorContext. + * To create multiple constraint violations with different variable values, this method can be called + * between successive calls to [[TwitterConstraintValidatorContext#addConstraintViolation]]. E.g., + * within a ConstraintValidator instance: + * {{{ + * def isValid(value: String, constraintValidatorContext: TwitterConstraintValidatorContext): Boolean = { + * + * TwitterConstraintValidatorContext + * .addMessageParameter("foo", "bar") + * .withMessageTemplate("{foo}") + * .addConstraintViolation(context) + * + * TwitterConstraintValidatorContext + * .addMessageParameter("foo", "snafu") + * .withMessageTemplate("{foo}") + * .addConstraintViolation(context) + * + * false + * }}} + * + * @param name the name under which to bind the parameter, cannot be `null`. + * @param value the value to be bound to the specified name + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def addMessageParameter(name: String, value: AnyRef): TwitterConstraintValidatorContext.Builder = + builder.addMessageParameter(name, value) + + /** + * Allows setting of an additional expression variable which will be available as an EL variable during interpolation. The + * variable will be available for interpolation for all constraint violations generated for this constraint. + * This includes the default one as well as all violations created by this TwitterConstraintValidatorContext. + * To create multiple constraint violations with different variable values, this method can be called + * between successive calls to [[TwitterConstraintValidatorContext#addConstraintViolation]]. E.g., + * within a ConstraintValidator instance: + * {{{ + * def isValid(value: String, constraintValidatorContext: TwitterConstraintValidatorContext): Boolean = { + * + * TwitterConstraintValidatorContext + * .addExpressionVariable("foo", "bar") + * .withMessageTemplate("${foo}") + * .addConstraintViolation(context) + * + * TwitterConstraintValidatorContext + * .addExpressionVariable("foo", "snafu") + * .withMessageTemplate("${foo}") + * .addConstraintViolation(context) + * + * false + * }}} + * + * @param name the name under which to bind the expression variable, cannot be `null`. + * @param value the value to be bound to the specified name + * + * @return a [[TwitterConstraintValidatorContext.Builder]] to allow for method chaining. + */ + def addExpressionVariable( + name: String, + value: AnyRef + ): TwitterConstraintValidatorContext.Builder = + builder.addExpressionVariable(name, value) +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/conversions/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/conversions/BUILD new file mode 100644 index 0000000000..da9927dad0 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/conversions/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-conversions", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-validator/src/main/scala/com/twitter/util/validation/engine", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/conversions/ConstraintViolationOps.scala b/util-validator/src/main/scala/com/twitter/util/validation/conversions/ConstraintViolationOps.scala new file mode 100644 index 0000000000..00c71e8e25 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/conversions/ConstraintViolationOps.scala @@ -0,0 +1,22 @@ +package com.twitter.util.validation.conversions + +import com.twitter.util.validation.engine.ConstraintViolationHelper +import jakarta.validation.{ConstraintViolation, Payload} + +object ConstraintViolationOps { + + implicit class RichConstraintViolation[T](val self: ConstraintViolation[T]) extends AnyVal { + + /** + * Returns a String created by concatenating the [[ConstraintViolation#getPropertyPath]] and + * the [[ConstraintViolation#getMessage]] separated by a colon (:). E.g. if a given violation + * has a `getPropertyPath.toString` value of "foo" and a `getMessage` value of "does not exist", + * this would return "foo: does not exist". + */ + def getMessageWithPath: String = ConstraintViolationHelper.messageWithPath(self) + + /** Lookup the dynamic payload of the given class type */ + def getDynamicPayload[P <: Payload](clazz: Class[P]): Option[P] = + ConstraintViolationHelper.dynamicPayload[P](self, clazz) + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/conversions/PathOps.scala b/util-validator/src/main/scala/com/twitter/util/validation/conversions/PathOps.scala new file mode 100644 index 0000000000..534c67c1a8 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/conversions/PathOps.scala @@ -0,0 +1,20 @@ +package com.twitter.util.validation.conversions + +import jakarta.validation.Path +import org.hibernate.validator.internal.engine.path.PathImpl + +object PathOps { + + implicit class RichPath(val self: Path) extends AnyVal { + def getLeafNode: Path.Node = self match { + case pathImpl: PathImpl => + pathImpl.getLeafNode + case _ => + var node = self.iterator().next() + while (self.iterator().hasNext) { + node = self.iterator().next() + } + node + } + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/engine/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/engine/BUILD new file mode 100644 index 0000000000..1b821957f3 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/engine/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-engine", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-slf4j-api/src/main/scala", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/engine/ConstraintViolationHelper.scala b/util-validator/src/main/scala/com/twitter/util/validation/engine/ConstraintViolationHelper.scala new file mode 100644 index 0000000000..55d4f68a9e --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/engine/ConstraintViolationHelper.scala @@ -0,0 +1,41 @@ +package com.twitter.util.validation.engine + +import com.twitter.util.{Return, Try} +import com.twitter.util.validation.internal.engine.ConstraintViolationFactory +import jakarta.validation.{ConstraintViolation, Payload} +import org.hibernate.validator.engine.HibernateConstraintViolation + +/** Helper for [[jakarta.validation.ConstraintViolation]] */ +object ConstraintViolationHelper { + + /** + * Sorts the given [[Set]] by natural order of String created by concatenating the + * [[ConstraintViolation#getPropertyPath]] and [[ConstraintViolation#getMessage]] separated + * by a colon (:). + * + * @see [[messageWithPath]] + */ + def sortViolations[T](violations: Set[ConstraintViolation[T]]): Seq[ConstraintViolation[T]] = + ConstraintViolationFactory.sortSet[T](violations) + + /** + * Returns a String created by concatenating the [[ConstraintViolation#getPropertyPath]] and + * the [[ConstraintViolation#getMessage]] separated by a colon (:). E.g. if a given violation + * has a `getPropertyPath.toString` value of "foo" and a `getMessage` value of "does not exist", + * this would return "foo: does not exist". + */ + def messageWithPath(violation: ConstraintViolation[_]): String = + ConstraintViolationFactory.getMessage(violation) + + /** Lookup the dynamic payload of the given class type */ + def dynamicPayload[P <: Payload]( + violation: ConstraintViolation[_], + clazz: Class[P] + ): Option[P] = { + Try(violation.unwrap(classOf[HibernateConstraintViolation[_]])) match { + case Return(hibernateConstraintViolation) => + Option(hibernateConstraintViolation.getDynamicPayload[P](clazz)) + case _ => None + } + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/engine/MethodValidationResult.scala b/util-validator/src/main/scala/com/twitter/util/validation/engine/MethodValidationResult.scala new file mode 100644 index 0000000000..9fd4ac8033 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/engine/MethodValidationResult.scala @@ -0,0 +1,88 @@ +package com.twitter.util.validation.engine + +import com.twitter.util.logging.Logger +import com.twitter.util.{Return, Throw, Try} +import jakarta.validation.Payload + +sealed trait MethodValidationResult { + def payload: Option[Payload] + def isValid: Boolean +} + +object MethodValidationResult { + private[this] val logger: Logger = Logger(MethodValidationResult.getClass) + + case object EmptyPayload extends Payload + + case object Valid extends MethodValidationResult { + override final val payload: Option[Payload] = None + override final val isValid: Boolean = true + } + + case class Invalid( + message: String, + payload: Option[Payload] = None) + extends MethodValidationResult { + override final val isValid: Boolean = false + } + + /** + * Utility for evaluating a condition in order to return a [[MethodValidationResult]]. Returns + * [[MethodValidationResult.Valid]] when the condition is `true`, otherwise if the condition + * evaluates to `false` or throws an exception a [[MethodValidationResult.Invalid]] will be returned. + * In the case of an exception, the `exception.getMessage` is used in place of the given message. + * + * @note This will not allow a non-fatal exception to escape. Instead a [[MethodValidationResult.Invalid]] + * will be returned when a non-fatal exception is encountered when evaluating `condition`. As + * this equates failure to execute the condition function to a return of `false`. + * @param condition function to evaluate for validation. + * @param message function to evaluate for a message when the given condition is `false`. + * @param payload [[Payload]] to use for when the given condition is `false`. + * + * @return a [[MethodValidationResult.Valid]] when the condition is `true` otherwise a [[MethodValidationResult.Invalid]]. + */ + def validIfTrue( + condition: => Boolean, + message: => String, + payload: Option[Payload] = None, + ): MethodValidationResult = Try(condition) match { + case Return(result) if result => + Valid + case Return(result) if !result => + Invalid(message, payload) + case Throw(e) => + logger.warn(e.getMessage, e) + Invalid(e.getMessage, payload) + } + + /** + * Utility for evaluating the negation of a condition in order to return a [[MethodValidationResult]]. + * Returns [[MethodValidationResult.Valid]] when the condition is `false`, otherwise if the condition + * evaluates to `true` or throws an exception a [[MethodValidationResult.Invalid]] will be returned. + * In the case of an exception, the `exception.getMessage` is used in place of the given message. + * + * @note This will not allow a non-fatal exception to escape. Instead a [[MethodValidationResult.Valid]] + * will be returned when a non-fatal exception is encountered when evaluating `condition`. As + * this equates failure to execute the condition to a return of `false`. + * @param condition function to evaluate for validation. + * @param message function to evaluate for a message when the given condition is `true`. + * @param payload [[Payload]] to use for when the given condition is `true`. + * + * @return a [[MethodValidationResult.Valid]] when the condition is `false` or when the condition evaluation + * throws a NonFatal exception otherwise a [[MethodValidationResult.Invalid]]. + */ + def validIfFalse( + condition: => Boolean, + message: => String, + payload: Option[Payload] = None + ): MethodValidationResult = + Try(condition) match { + case Return(result) if !result => + Valid + case Return(result) if result => + Invalid(message, payload) + case Throw(e) => + logger.warn(e.getMessage, e) + Valid + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/executable/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/executable/BUILD new file mode 100644 index 0000000000..95829dce5a --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/executable/BUILD @@ -0,0 +1,20 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-executable", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/executable/ScalaExecutableValidator.scala b/util-validator/src/main/scala/com/twitter/util/validation/executable/ScalaExecutableValidator.scala new file mode 100644 index 0000000000..62f77809fc --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/executable/ScalaExecutableValidator.scala @@ -0,0 +1,194 @@ +package com.twitter.util.validation.executable + +import com.twitter.util.validation.metadata.{ExecutableDescriptor, MethodDescriptor} +import jakarta.validation.{ConstraintViolation, ValidationException} +import java.lang.reflect.{Constructor, Method} + +/** + * Scala version of the Bean Specification [[jakarta.validation.executable.ExecutableValidator]]. + * + * Validates parameters and return values of methods and constructors. Implementations of this interface must be thread-safe. + */ +trait ScalaExecutableValidator { + + /** + * Validates all methods annotated with `@MethodValidation` of the given object. + * + * @param obj the object on which the methods to validate are invoked. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the methods to validate. + * + * @return a set with the constraint violations caused by this validation; will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for the object to validate. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateMethods[T]( + obj: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates the given method annotated with `@MethodValidation` of the given object. + * + * @param obj the object on which the method to validate is invoked. + * @param method the `@MethodValidation` annotated method to invoke for validation. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the method to validate. + * + * @return a set with the constraint violations caused by this validation; will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for the object to validate or if the given method is not a method of the object class type. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateMethod[T]( + obj: T, + method: Method, + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all constraints placed on the parameters of the given method. + * + * @param obj the object on which the method to validate is invoked. + * @param method the method for which the parameter constraints is validated. + * @param parameterValues the values provided by the caller for the given method's parameters. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the method to validate. + * + * @return a set with the constraint violations caused by this validation; will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateParameters[T]( + obj: T, + method: Method, + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all return value constraints of the given method. + * + * @param obj the object on which the method to validate is invoked. + * @param method the method for which the return value constraints is validated. + * @param returnValue the value returned by the given method. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the method to validate. + * + * @return a set with the constraint violations caused by this validation; will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateReturnValue[T]( + obj: T, + method: Method, + returnValue: Any, + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all constraints placed on the parameters of the given constructor. + * + * @param constructor the constructor for which the parameter constraints is validated. + * @param parameterValues the values provided by the caller for the given constructor's parameters. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the constructor to validate. + * + * @return a set with the constraint violations caused by this validation; Will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateConstructorParameters[T]( + constructor: Constructor[T], + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all constraints placed on the parameters of the given method. + * + * @param method the method for which the parameter constraints is validated. + * @param parameterValues the values provided by the caller for the given method's parameters. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type defining the method to validate. + * + * @return a set with the constraint violations caused by this validation; Will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateMethodParameters[T]( + method: Method, + parameterValues: Array[Any], + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all return value constraints of the given constructor. + * + * @param constructor the constructor for which the return value constraints is validated. + * @param createdObject the object instantiated by the given method. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type hosting the constructor to validate. + * + * @return a set with the constraint violations caused by this validation; will be empty, if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non recoverable error happens during the validation process. + */ + @throws[ValidationException] + def validateConstructorReturnValue[T]( + constructor: Constructor[T], + createdObject: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates all constraints placed on the parameters of the given executable. + * + * @param executable the [[ExecutableDescriptor]] for which the parameter constraints is validated. + * @param parameterValues the values provided by the caller for the given executable's parameters. + * @param parameterNames the parameter names to use for error reporting. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type defining the executable to validate. + * + * @return a set with the constraint violations caused by this validation; Will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non-recoverable error happens during the validation process. + */ + private[twitter] def validateExecutableParameters[T]( + executable: ExecutableDescriptor, + parameterValues: Array[Any], + parameterNames: Array[String], + groups: Class[_]* + ): Set[ConstraintViolation[T]] + + /** + * Validates the `@MethodValidation` placed constraints on the given array of descriptors. + * @param methods the array of [[MethodDescriptor]] instances to evaluate. + * @param obj the object on which the method to validate is invoked. + * @param groups the list of groups targeted for validation (defaults to Default). + * @tparam T the type defining the methods to validate. + * + * @return a set with the constraint violations caused by this validation; Will be empty if no error occurs, but never null. + * + * @throws IllegalArgumentException - if null is passed for any of the parameters or if parameters don't match with each other. + * @throws ValidationException - if a non-recoverable error happens during the validation process. + */ + private[twitter] def validateMethods[T]( + methods: Array[MethodDescriptor], + obj: T, + groups: Class[_]* + ): Set[ConstraintViolation[T]] +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/AnnotationFactory.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/AnnotationFactory.scala new file mode 100644 index 0000000000..ede256adc0 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/AnnotationFactory.scala @@ -0,0 +1,48 @@ +package com.twitter.util.validation.internal + +import org.hibernate.validator.internal.util.annotation.AnnotationDescriptor + +import java.lang.annotation.Annotation +import java.lang.reflect.Method +import java.util + +private[validation] object AnnotationFactory { + + /** + * Creates a new Annotation (AnnotationProxy) instance of + * the given annotation type with the custom values applied. + * + */ + def newInstance[A <: Annotation]( + annotationType: Class[A], + customValues: Map[String, Any] = Map.empty + ): A = { + val attributes: util.Map[String, Object] = new util.HashMap[String, Object]() + // Extract default values from annotation + val methods: Array[Method] = annotationType.getDeclaredMethods + var index = 0 + val length = methods.length + while (index < length) { + val method = methods(index) + attributes.put(method.getName, method.getDefaultValue) + index += 1 + } + // Populate custom values + val keysIterator = customValues.keysIterator + while (keysIterator.hasNext) { + val key = keysIterator.next() + attributes.put(key, customValues(key).asInstanceOf[AnyRef]) + } + annotationForMap(annotationType, attributes).asInstanceOf[A] + } + + private def annotationForMap[A <: Annotation]( + annotationClass: Class[A], + memberValues: util.Map[String, AnyRef] + ): Annotation = { + val ad: AnnotationDescriptor[A] = + new AnnotationDescriptor.Builder[A](annotationClass, memberValues).build() + org.hibernate.validator.internal.util.annotation.AnnotationFactory + .create(ad) + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/BUILD new file mode 100644 index 0000000000..577f8846a6 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-core:scala", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/Types.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/Types.scala new file mode 100644 index 0000000000..42eae7840c --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/Types.scala @@ -0,0 +1,84 @@ +package com.twitter.util.validation.internal + +import org.json4s.reflect.Reflector +import org.json4s.reflect.ScalaType + +private[validation] object Types { + + /** For certain scala types, return the scala type of the first type arg */ + def getContainedScalaType( + scalaType: ScalaType + ): Option[ScalaType] = scalaType match { + case argType if argType.isMap || argType.isMutableMap => + // Maps not supported + None + case argType if argType.isCollection || argType.isOption => + // handles Option and Iterable collections + if (argType.typeArgs.isEmpty) Some(Reflector.scalaTypeOf(classOf[Object])) + else Some(argType.typeArgs.head) + case _ => + Some(scalaType) + } + + // if the case class is generically typed, the originally parsed scala type + // may be an object thus we try to refine it based on the given value + def refineScalaType(value: Any, scalaType: ScalaType): ScalaType = { + if (scalaType.erasure == classOf[Object] && value != null) { + if (classOf[Option[_]].isAssignableFrom(value.getClass)) + Reflector.scalaTypeOf(classOf[Option[_]]) + else + Reflector.scalaTypeOf(value.getClass) + } else if (scalaType.isOption) { + if (scalaType.typeArgs.nonEmpty) { + // When deserializing a case class, we interpret the field + // type with the field constructor type. For optional + // field, Jackson returns Option[Object] as the constructor + // type. However, field validations would require a + // specific field type in order to run the constraint + // validator. So we use the field value type as the field + // type for optional field. + if (scalaType.typeArgs.head.erasure.isInstanceOf[Object]) + Reflector.scalaTypeOf(value.getClass) + else scalaType.typeArgs.head + } else { + Reflector.scalaTypeOf(classOf[Object]) + } + } else scalaType + } + + /** Translate a Scala type to a corresponding Java class */ + def getJavaType(scalaType: ScalaType): Class[_] = scalaType match { + case argType if argType.isOption => + if (argType.typeArgs.nonEmpty) argType.typeArgs.head.erasure + else classOf[Object] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Byte.TYPE.getSimpleName => + classOf[java.lang.Byte] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Short.TYPE.getSimpleName => + classOf[java.lang.Short] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Character.TYPE.getSimpleName => + classOf[java.lang.Character] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Integer.TYPE.getSimpleName => + classOf[java.lang.Integer] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Long.TYPE.getSimpleName => + classOf[java.lang.Long] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Float.TYPE.getSimpleName => + classOf[java.lang.Float] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Double.TYPE.getSimpleName => + classOf[java.lang.Double] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Boolean.TYPE.getSimpleName => + classOf[java.lang.Boolean] + case argType + if argType.isPrimitive && argType.simpleName == java.lang.Void.TYPE.getSimpleName => + classOf[java.lang.Void] + case argType => + argType.erasure + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/ValidationContext.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/ValidationContext.scala new file mode 100644 index 0000000000..f730837c08 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/ValidationContext.scala @@ -0,0 +1,12 @@ +package com.twitter.util.validation.internal + +import org.hibernate.validator.internal.engine.path.PathImpl + +/** Utility class to carry necessary context for validation */ +private[validation] case class ValidationContext[T]( + fieldName: Option[String], + rootClazz: Option[Class[T]], + root: Option[T], + leaf: Option[Any], + path: PathImpl, + isFailFast: Boolean = false) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/BUILD new file mode 100644 index 0000000000..90be7adb9f --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/BUILD @@ -0,0 +1,17 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-constraintvalidation", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-validator/src/main/scala/org/hibernate/validator/internal/engine", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/ConstraintValidatorContextFactory.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/ConstraintValidatorContextFactory.scala new file mode 100644 index 0000000000..9227b6fd4a --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/constraintvalidation/ConstraintValidatorContextFactory.scala @@ -0,0 +1,26 @@ +package com.twitter.util.validation.internal.constraintvalidation + +import jakarta.validation.ConstraintValidatorContext +import jakarta.validation.metadata.ConstraintDescriptor +import org.hibernate.validator.internal.engine.ValidatorFactoryInspector +import org.hibernate.validator.internal.engine.constraintvalidation.{ + ConstraintValidatorContextImpl => HibernateConstraintValidatorContextImpl +} +import org.hibernate.validator.internal.engine.path.PathImpl + +private[validation] class ConstraintValidatorContextFactory( + validatorFactory: ValidatorFactoryInspector) { + + def newConstraintValidatorContext( + path: PathImpl, + constraintDescriptor: ConstraintDescriptor[_] + ): ConstraintValidatorContext = + new HibernateConstraintValidatorContextImpl( + validatorFactory.validatorFactoryScopedContext.getClockProvider, + path, + constraintDescriptor, + validatorFactory.validatorFactoryScopedContext.getConstraintValidatorPayload, + validatorFactory.validatorFactoryScopedContext.getConstraintExpressionLanguageFeatureLevel, + validatorFactory.validatorFactoryScopedContext.getCustomViolationExpressionLanguageFeatureLevel + ) +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/BUILD new file mode 100644 index 0000000000..d42615e719 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/BUILD @@ -0,0 +1,20 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-engine", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/scala-lang/modules:scala-collection-compat", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor", + "util/util-validator/src/main/scala/org/hibernate/validator/internal/engine", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/ConstraintViolationFactory.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/ConstraintViolationFactory.scala new file mode 100644 index 0000000000..e59da4626d --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/engine/ConstraintViolationFactory.scala @@ -0,0 +1,167 @@ +package com.twitter.util.validation.internal.engine + +import jakarta.validation.metadata.ConstraintDescriptor +import jakarta.validation.{ + ConstraintValidatorContext, + ConstraintViolation, + ConstraintViolationException, + Payload +} +import java.util.Collections +import org.hibernate.validator.internal.engine.constraintvalidation.{ + ConstraintValidatorContextImpl, + ConstraintViolationCreationContext +} +import org.hibernate.validator.internal.engine.path.PathImpl +import org.hibernate.validator.internal.engine.{ + MessageInterpolatorContext, + ValidatorFactoryInspector, + ConstraintViolationImpl => HibernateConstraintViolationImpl +} +import scala.collection.{SortedSet, mutable} +import scala.jdk.CollectionConverters._ + +private[validation] object ConstraintViolationFactory { + private[this] val MessageWithPathTemplate = "%s: %s" + + def sortSet[T](unsortedSet: Set[ConstraintViolation[T]]): Seq[ConstraintViolation[T]] = + (SortedSet.empty[ConstraintViolation[T]](new ConstraintViolationOrdering) ++ unsortedSet).toSeq + + /* Private */ + + def getMessage(violation: ConstraintViolation[_]): String = + MessageWithPathTemplate.format(violation.getPropertyPath.toString, violation.getMessage) + + private[this] class ConstraintViolationOrdering[T] extends Ordering[ConstraintViolation[T]] { + override def compare(x: ConstraintViolation[T], y: ConstraintViolation[T]): Int = + getMessage(x).compare(getMessage(y)) + } +} + +/** A wrapper/helper for [[ConstraintViolation]] instances */ +private[validation] class ConstraintViolationFactory(validatorFactory: ValidatorFactoryInspector) { + import ConstraintViolationFactory._ + + /** Instantiate a new [[ConstraintViolationException]] from the given Set of violations */ + def newConstraintViolationException( + violations: Set[ConstraintViolation[Any]] + ): ConstraintViolationException = { + val errors: Seq[ConstraintViolation[Any]] = sortSet(violations) + + val message: String = + "\nValidation Errors:\t\t" + errors + .map(getMessage).mkString(", ") + "\n\n" + new ConstraintViolationException(message, errors.toSet.asJava) + } + + /** + * Performs message interpolation given the constraint descriptor and constraint validator context + * to create a set of [[ConstraintViolation]] from the given context and parameters. + */ + def buildConstraintViolations[T]( + rootClazz: Option[Class[T]], + root: Option[T], + leaf: Option[Any], + path: PathImpl, + invalidValue: Any, + constraintDescriptor: ConstraintDescriptor[_], + constraintValidatorContext: ConstraintValidatorContext + ): Set[ConstraintViolation[T]] = { + val results = mutable.HashSet[ConstraintViolation[T]]() + val constraintViolationCreationContexts = + constraintValidatorContext + .asInstanceOf[ConstraintValidatorContextImpl] + .getConstraintViolationCreationContexts + + var index = 0 + val size = constraintViolationCreationContexts.size() + while (index < size) { + val constraintViolationCreationContext = constraintViolationCreationContexts.get(index) + val messageTemplate = constraintViolationCreationContext.getMessage + val interpolatedMessage = + validatorFactory.messageInterpolator + .interpolate( + messageTemplate, + new MessageInterpolatorContext( + constraintDescriptor, + invalidValue, + root.map(_.getClass.asInstanceOf[Class[T]]).orNull, + constraintViolationCreationContext.getPath, + constraintViolationCreationContext.getMessageParameters, + constraintViolationCreationContext.getExpressionVariables, + constraintViolationCreationContext.getExpressionLanguageFeatureLevel, + constraintViolationCreationContext.isCustomViolation + ) + ) + results.add( + newConstraintViolation[T]( + messageTemplate, + interpolatedMessage, + path, + invalidValue, + rootClazz.orNull, + root.map(_.asInstanceOf[T]).getOrElse(null.asInstanceOf[T]), + leaf.orNull, + constraintDescriptor, + constraintViolationCreationContext + ) + ) + index += 1 + } + + if (results.nonEmpty) results.toSet + else Set.empty[ConstraintViolation[T]] + } + + /** Creates a new [[ConstraintViolation]] from the given context and parameters. */ + def newConstraintViolation[T]( + messageTemplate: String, + interpolatedMessage: String, + path: PathImpl, + invalidValue: Any, + rootClazz: Class[T], + root: T, + leaf: Any, + constraintDescriptor: ConstraintDescriptor[_], + payload: Payload + ): ConstraintViolation[T] = + HibernateConstraintViolationImpl.forBeanValidation[T]( + messageTemplate, + Collections.emptyMap[String, Object](), + Collections.emptyMap[String, Object](), + interpolatedMessage, + rootClazz, + root, + leaf, + invalidValue, + path, + constraintDescriptor, + payload + ) + + /** Creates a new [[ConstraintViolation]] from the given context and parameters. */ + def newConstraintViolation[T]( + messageTemplate: String, + interpolatedMessage: String, + path: PathImpl, + invalidValue: Any, + rootClazz: Class[T], + root: T, + leaf: Any, + constraintDescriptor: ConstraintDescriptor[_], + constraintViolationCreationContext: ConstraintViolationCreationContext + ): ConstraintViolation[T] = + HibernateConstraintViolationImpl.forBeanValidation[T]( + messageTemplate, + constraintViolationCreationContext.getMessageParameters, + constraintViolationCreationContext.getExpressionVariables, + interpolatedMessage, + rootClazz, + root, + leaf, + invalidValue, + path, + constraintDescriptor, + constraintViolationCreationContext.getDynamicPayload + ) +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/BUILD new file mode 100644 index 0000000000..ff9bf3312c --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/BUILD @@ -0,0 +1,18 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-metadata-descriptor", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation", + "util/util-validator/src/main/scala/org/hibernate/validator/internal/engine", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/ConstraintDescriptorFactory.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/ConstraintDescriptorFactory.scala new file mode 100644 index 0000000000..e757cc9c0d --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/ConstraintDescriptorFactory.scala @@ -0,0 +1,52 @@ +package com.twitter.util.validation.internal.metadata.descriptor + +import jakarta.validation.groups.Default +import java.lang.annotation.Annotation +import java.lang.reflect.Type +import org.hibernate.validator.internal.engine.ValidatorFactoryInspector +import org.hibernate.validator.internal.metadata.core.ConstraintOrigin +import org.hibernate.validator.internal.metadata.descriptor.{ + ConstraintDescriptorImpl => HibernateConstraintDescriptorImpl +} +import org.hibernate.validator.internal.metadata.descriptor.ConstraintDescriptorImpl.ConstraintType +import org.hibernate.validator.internal.metadata.location.ConstraintLocation.ConstraintLocationKind +import org.hibernate.validator.internal.metadata.raw.ConstrainedElement +import org.hibernate.validator.internal.metadata.raw.ConstrainedElement.ConstrainedElementKind +import org.hibernate.validator.internal.properties.Constrainable +import org.hibernate.validator.internal.util.annotation.ConstraintAnnotationDescriptor + +private[validation] class ConstraintDescriptorFactory( + validatorFactory: ValidatorFactoryInspector) { + + def newConstraintDescriptor( + name: String, + clazz: Class[_], + declaringClazz: Class[_], + annotation: Annotation, + constrainedElementKind: ConstrainedElementKind = ConstrainedElementKind.FIELD + ) = new HibernateConstraintDescriptorImpl( + validatorFactory.constraintHelper, + mkConstrainable(name, clazz, declaringClazz, constrainedElementKind), + new ConstraintAnnotationDescriptor.Builder(annotation).build(), + ConstraintLocationKind.FIELD, + classOf[Default], + ConstraintOrigin.DEFINED_LOCALLY, + ConstraintType.GENERIC + ) + + /* Private */ + + private[this] def mkConstrainable( + name: String, + clazz: Class[_], + declaringClazz: Class[_], + constrainedElementKind: ConstrainedElementKind = ConstrainedElementKind.FIELD + ): Constrainable = new Constrainable { + override def getName: String = name + override def getDeclaringClass: Class[_] = declaringClazz + override def getTypeForValidatorResolution: Type = clazz + override def getType: Type = clazz + override def getConstrainedElementKind: ConstrainedElement.ConstrainedElementKind = + constrainedElementKind + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/MethodValidationConstraintDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/MethodValidationConstraintDescriptor.scala new file mode 100644 index 0000000000..11e6b27df9 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/metadata/descriptor/MethodValidationConstraintDescriptor.scala @@ -0,0 +1,41 @@ +package com.twitter.util.validation.internal.metadata.descriptor + +import com.twitter.util.validation.MethodValidation +import jakarta.validation.metadata.{ConstraintDescriptor, ValidateUnwrappedValue} +import jakarta.validation.{ConstraintTarget, ConstraintValidator, Payload, ValidationException} +import java.util + +private[validation] class MethodValidationConstraintDescriptor(annotation: MethodValidation) + extends ConstraintDescriptor[MethodValidation] { + override def getAnnotation: MethodValidation = annotation + override def getMessageTemplate: String = annotation.message() + override def getGroups: util.Set[Class[_]] = { + val target = new util.HashSet[Class[_]] + annotation.groups().foreach(target.add) + target + } + override def getPayload: util.Set[Class[_ <: Payload]] = { + val target = new util.HashSet[Class[_ <: Payload]] + annotation.payload().foreach(target.add) + target + } + override def getValidationAppliesTo: ConstraintTarget = ConstraintTarget.PARAMETERS + override def getConstraintValidatorClasses: util.List[Class[ + _ <: ConstraintValidator[MethodValidation, _] + ]] = util.Collections.emptyList() + override def getAttributes: util.Map[String, AnyRef] = { + val target = new util.HashMap[String, AnyRef]() + var i = 0 + annotation.fields().foreach { key => + target.put(key, Integer.valueOf(i)) + i += 1 + } + target + } + override def getComposingConstraints: util.Set[ConstraintDescriptor[_]] = + util.Collections.emptySet() + override def isReportAsSingleViolation: Boolean = true + override def getValueUnwrapping: ValidateUnwrappedValue = ValidateUnwrappedValue.DEFAULT + override def unwrap[U](clazz: Class[U]): U = throw new ValidationException( + s"${clazz.getName} is unsupported") +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/BUILD new file mode 100644 index 0000000000..dd9baf6944 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/BUILD @@ -0,0 +1,20 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-validators", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-slf4j-api/src/main/scala/com/twitter/util/logging", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints", + "util/util-validator/src/main/scala/com/twitter/util/validation/constraintvalidation", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidator.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidator.scala new file mode 100644 index 0000000000..bd3c56270b --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidator.scala @@ -0,0 +1,63 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.constraints.CountryCode +import com.twitter.util.validation.constraintvalidation.TwitterConstraintValidatorContext +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import java.util.Locale + +private object ISO3166CountryCodeConstraintValidator { + + /** @see [[https://www.iso.org/iso-3166-country-codes.html ISO 3166]] */ + val CountryCodes: Set[String] = Locale.getISOCountries.toSet +} + +private[validation] class ISO3166CountryCodeConstraintValidator + extends ConstraintValidator[CountryCode, Any] { + import ISO3166CountryCodeConstraintValidator._ + + @volatile private[this] var countryCode: CountryCode = _ + + override def initialize(constraintAnnotation: CountryCode): Unit = { + this.countryCode = constraintAnnotation + } + + override def isValid( + obj: Any, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = obj match { + case null => true + case typedValue: Array[Any] => + validationResult(typedValue, constraintValidatorContext) + case typedValue: Iterable[Any] => + validationResult(typedValue, constraintValidatorContext) + case anyValue => + validationResult(Seq(anyValue.toString), constraintValidatorContext) + } + + /* Private */ + + private[this] def validationResult( + value: Iterable[Any], + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + val invalidCountryCodes = findInvalidCountryCodes(value) + // an empty value is not a valid country code + val valid = if (value.isEmpty) false else invalidCountryCodes.isEmpty + + if (!valid) { + TwitterConstraintValidatorContext + .addExpressionVariable("validatedValue", mkString(value)) + .withMessageTemplate(countryCode.message()) + .addConstraintViolation(constraintValidatorContext) + } + + valid + } + + private[this] def findInvalidCountryCodes(values: Iterable[Any]): Set[String] = { + val uppercaseCountryCodes = values.toSet.map { value: Any => + value.toString.toUpperCase + } + uppercaseCountryCodes.diff(CountryCodes) + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidator.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidator.scala new file mode 100644 index 0000000000..9f2bea72e4 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidator.scala @@ -0,0 +1,57 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.constraints.OneOf +import com.twitter.util.validation.constraintvalidation.TwitterConstraintValidatorContext +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + +private[validation] class OneOfConstraintValidator extends ConstraintValidator[OneOf, Any] { + + @volatile private[this] var oneOf: OneOf = _ + @volatile private[this] var values: Set[String] = _ + + override def initialize(constraintAnnotation: OneOf): Unit = { + this.oneOf = constraintAnnotation + this.values = constraintAnnotation.value().toSet + } + + override def isValid( + obj: Any, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = obj match { + case null => true + case arrayValue: Array[_] => + validationResult(arrayValue, values, constraintValidatorContext) + case traversableValue: Iterable[_] => + validationResult(traversableValue, values, constraintValidatorContext) + case anyValue => + validationResult(Seq(anyValue.toString), values, constraintValidatorContext) + } + + /* Private */ + + private[this] def validationResult( + value: Iterable[_], + oneOfValues: Set[String], + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + val invalidValues = findInvalidValues(value, oneOfValues) + // an empty value is not one of the given values + val valid = if (value.isEmpty) false else invalidValues.isEmpty + if (!valid) { + TwitterConstraintValidatorContext + .addExpressionVariable("validatedValue", mkString(value)) + .withMessageTemplate(oneOf.message()) + .addConstraintViolation(constraintValidatorContext) + } + + valid + } + + private[this] def findInvalidValues( + value: Iterable[_], + oneOfValues: Set[String] + ): Set[String] = { + val valueAsStrings = value.map(_.toString).toSet + valueAsStrings.diff(oneOfValues) + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidator.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidator.scala new file mode 100644 index 0000000000..2220a6facd --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidator.scala @@ -0,0 +1,17 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.Try +import com.twitter.util.validation.constraints.UUID +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import java.util.{UUID => JUUID} + +private[validation] class UUIDConstraintValidator extends ConstraintValidator[UUID, String] { + + override def isValid( + obj: String, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + if (obj == null) true + else Try(JUUID.fromString(obj)).isReturn + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/BUILD new file mode 100644 index 0000000000..86fd42e9a9 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/BUILD @@ -0,0 +1,18 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-validators-notempty", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/NotEmptyValidatorForIterable.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/NotEmptyValidatorForIterable.scala new file mode 100644 index 0000000000..cb00473018 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/notempty/NotEmptyValidatorForIterable.scala @@ -0,0 +1,23 @@ +package com.twitter.util.validation.internal.validators.notempty + +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import jakarta.validation.constraints.NotEmpty + +/** + * Intended to work exactly like [[org.hibernate.validator.internal.constraintvalidators.bv.notempty.NotEmptyValidatorForArray]] + * but for Scala [[Iterable]] types. + * + * Check that the [[Iterable]] is not null and not empty. + * + * @see [[org.hibernate.validator.internal.constraintvalidators.bv.notempty.NotEmptyValidatorForArray]] + */ +private[validation] class NotEmptyValidatorForIterable + extends ConstraintValidator[NotEmpty, Iterable[_]] { + override def isValid( + obj: Iterable[_], + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + if (obj == null) false + else obj.nonEmpty + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/package.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/package.scala new file mode 100644 index 0000000000..0dee50a7ea --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/package.scala @@ -0,0 +1,11 @@ +package com.twitter.util.validation.internal + +package object validators { + + private[validation] def mkString(value: Iterable[_]): String = { + val trimmed = value.map(_.toString.trim).filter(_.nonEmpty) + if (trimmed.nonEmpty && trimmed.size > 1) s"${trimmed.mkString("[", ", ", "]")}" + else if (trimmed.nonEmpty) trimmed.head + else "" + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/BUILD new file mode 100644 index 0000000000..2af10aac5d --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/BUILD @@ -0,0 +1,18 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-validators-size", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-core:scala", + "util/util-validator-constraints/src/main/java/com/twitter/util/validation/constraints", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/SizeValidatorForIterable.scala b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/SizeValidatorForIterable.scala new file mode 100644 index 0000000000..b4d168b5e5 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/internal/validators/size/SizeValidatorForIterable.scala @@ -0,0 +1,38 @@ +package com.twitter.util.validation.internal.validators.size + +import jakarta.validation.constraints.Size +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + +/** + * Intended to work exactly like [[org.hibernate.validator.internal.constraintvalidators.bv.size.SizeValidatorForArray]] + * but for Scala [[Iterable]] types. + * + * Check that the length of an [[Iterable]] is between `min` and `max`. + * + * @see [[org.hibernate.validator.internal.constraintvalidators.bv.size.SizeValidatorForArray]] + */ +private[validation] class SizeValidatorForIterable extends ConstraintValidator[Size, Iterable[_]] { + + @volatile private[this] var min: Int = _ + @volatile private[this] var max: Int = _ + + override def initialize(constraintAnnotation: Size): Unit = { + this.min = constraintAnnotation.min() + this.max = constraintAnnotation.max() + validateParameters() + } + + override def isValid( + obj: Iterable[_], + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + if (obj == null) true + else obj.size >= min && obj.size <= max + } + + private def validateParameters(): Unit = { + if (min < 0) throw new IllegalArgumentException("The min parameter cannot be negative.") + if (max < 0) throw new IllegalArgumentException("The max parameter cannot be negative.") + if (max < min) throw new IllegalArgumentException("The length cannot be negative.") + } +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/BUILD b/util-validator/src/main/scala/com/twitter/util/validation/metadata/BUILD new file mode 100644 index 0000000000..eba48d8d0b --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-metadata", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/json4s:json4s-core", + "util/util-core:scala", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala", + ], +) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/CaseClassDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/CaseClassDescriptor.scala new file mode 100644 index 0000000000..8611c96423 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/CaseClassDescriptor.scala @@ -0,0 +1,13 @@ +package com.twitter.util.validation.metadata + +import java.lang.annotation.Annotation +import org.json4s.reflect.ScalaType + +case class CaseClassDescriptor[T] private[validation] ( + clazz: Class[T], + scalaType: ScalaType, + annotations: Array[Annotation], + constructors: Array[ConstructorDescriptor], + members: Map[String, PropertyDescriptor], + methods: Array[MethodDescriptor]) + extends Descriptor diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/ConstructorDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/ConstructorDescriptor.scala new file mode 100644 index 0000000000..d66c9fd29c --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/ConstructorDescriptor.scala @@ -0,0 +1,10 @@ +package com.twitter.util.validation.metadata + +import java.lang.annotation.Annotation +import java.lang.reflect.Executable + +case class ConstructorDescriptor private[validation] ( + override val executable: Executable, + annotations: Array[Annotation], + members: Map[String, PropertyDescriptor]) + extends ExecutableDescriptor(executable) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/Descriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/Descriptor.scala new file mode 100644 index 0000000000..7d6f72f03d --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/Descriptor.scala @@ -0,0 +1,4 @@ +package com.twitter.util.validation.metadata + +/** Marker trait for Descriptor Metadata */ +trait Descriptor diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/ExecutableDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/ExecutableDescriptor.scala new file mode 100644 index 0000000000..8157d4aac9 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/ExecutableDescriptor.scala @@ -0,0 +1,10 @@ +package com.twitter.util.validation.metadata + +import java.lang.annotation.Annotation +import java.lang.reflect.Executable + +abstract class ExecutableDescriptor private[validation] (val executable: Executable) + extends Descriptor { + def annotations: Array[Annotation] + def members: Map[String, PropertyDescriptor] +} diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/MethodDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/MethodDescriptor.scala new file mode 100644 index 0000000000..7484ea02d8 --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/MethodDescriptor.scala @@ -0,0 +1,10 @@ +package com.twitter.util.validation.metadata + +import java.lang.annotation.Annotation +import java.lang.reflect.Method + +case class MethodDescriptor private[validation] ( + method: Method, + annotations: Array[Annotation], + members: Map[String, PropertyDescriptor]) + extends ExecutableDescriptor(method) diff --git a/util-validator/src/main/scala/com/twitter/util/validation/metadata/PropertyDescriptor.scala b/util-validator/src/main/scala/com/twitter/util/validation/metadata/PropertyDescriptor.scala new file mode 100644 index 0000000000..6f94d7920f --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/metadata/PropertyDescriptor.scala @@ -0,0 +1,11 @@ +package com.twitter.util.validation.metadata + +import java.lang.annotation.Annotation +import org.json4s.reflect.ScalaType + +case class PropertyDescriptor private[validation] ( + scalaType: ScalaType, + cascadedScalaType: Option[ScalaType], + annotations: Array[Annotation], + isCascaded: Boolean = false) + extends Descriptor diff --git a/util-validator/src/main/scala/com/twitter/util/validation/package.scala b/util-validator/src/main/scala/com/twitter/util/validation/package.scala new file mode 100644 index 0000000000..d146cd8cde --- /dev/null +++ b/util-validator/src/main/scala/com/twitter/util/validation/package.scala @@ -0,0 +1,10 @@ +package com.twitter.util + +import jakarta.validation.ConstraintValidator +import java.lang.annotation.Annotation + +package object validation { + type ConstraintValidatorType = ConstraintValidator[_ <: Annotation, _] + type ConstraintAnnotationClazz = Class[_ <: Annotation] + type ConstraintValidatorClazz = Class[_ <: ConstraintValidatorType] +} diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/engine/BUILD b/util-validator/src/main/scala/org/hibernate/validator/internal/engine/BUILD new file mode 100644 index 0000000000..e5bff6803b --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/engine/BUILD @@ -0,0 +1,20 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-hibernate-engine", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + ], +) diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/engine/ValidatorFactoryInspector.scala b/util-validator/src/main/scala/org/hibernate/validator/internal/engine/ValidatorFactoryInspector.scala new file mode 100644 index 0000000000..1acde1ef9f --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/engine/ValidatorFactoryInspector.scala @@ -0,0 +1,41 @@ +package org.hibernate.validator.internal.engine + +import jakarta.validation.{ConstraintValidatorFactory, MessageInterpolator} +import java.lang.reflect.Field +import org.hibernate.validator.internal.metadata.core.ConstraintHelper +import org.hibernate.validator.internal.util.{ExecutableHelper, ExecutableParameterNameProvider} + +/** Allows access to the configured [[ConstraintHelper]] and [[ConstraintValidatorFactory]] */ +class ValidatorFactoryInspector(underlying: ValidatorFactoryImpl) { + + def constraintHelper: ConstraintHelper = + underlying.getConstraintCreationContext.getConstraintHelper + + def constraintValidatorFactory: ConstraintValidatorFactory = + underlying.getConstraintValidatorFactory + + def validatorFactoryScopedContext: ValidatorFactoryScopedContext = + underlying.getValidatorFactoryScopedContext + + def messageInterpolator: MessageInterpolator = + underlying.getMessageInterpolator + + def getConstraintCreationContext: ConstraintCreationContext = + underlying.getConstraintCreationContext + + def getExecutableHelper: ExecutableHelper = { + val field: Field = classOf[ValidatorFactoryImpl].getDeclaredField("executableHelper") + field.setAccessible(true) + field.get(underlying).asInstanceOf[ExecutableHelper] + } + + def getExecutableParameterNameProvider: ExecutableParameterNameProvider = + underlying.getExecutableParameterNameProvider + + def getMethodValidationConfiguration: MethodValidationConfiguration = + underlying.getMethodValidationConfiguration + + /** Close underlying ValidatorFactoryImpl */ + def close(): Unit = + underlying.close() +} diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/BUILD b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/BUILD new file mode 100644 index 0000000000..71e13675c0 --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/BUILD @@ -0,0 +1,21 @@ +scala_library( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-validation-internal-hibernate-properties-javabeans", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + ], +) diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ConstructorCallable.scala b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ConstructorCallable.scala new file mode 100644 index 0000000000..b52f66cad3 --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ConstructorCallable.scala @@ -0,0 +1,20 @@ +package org.hibernate.validator.internal.properties.javabean + +import com.twitter.util.reflect.Classes +import java.lang.reflect.{Constructor, Method, TypeVariable} +import org.hibernate.validator.internal.metadata.raw.ConstrainedElement.ConstrainedElementKind + +class ConstructorCallable(constructor: Constructor[_]) + extends ExecutableCallable[Constructor[_]](constructor, true) { + + override def getName: String = Classes.simpleName(getDeclaringClass) + + def getConstrainedElementKind: ConstrainedElementKind = ConstrainedElementKind.CONSTRUCTOR + + def getTypeParameters: Array[TypeVariable[Class[_]]] = + executable + .asInstanceOf[Method] + .getReturnType + .getTypeParameters + .asInstanceOf[Array[TypeVariable[Class[_]]]] +} diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ExecutableCallable.scala b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ExecutableCallable.scala new file mode 100644 index 0000000000..de74eac668 --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/ExecutableCallable.scala @@ -0,0 +1,186 @@ +package org.hibernate.validator.internal.properties.javabean + +import com.twitter.util.reflect.Classes +import jakarta.validation.ValidationException +import java.lang.reflect.{Constructor, Executable, Method, Modifier, Parameter, Type, TypeVariable} +import org.hibernate.validator.internal.properties.{Callable, Signature} +import org.hibernate.validator.internal.util.{ + ExecutableHelper, + ExecutableParameterNameProvider, + ReflectionHelper, + TypeHelper +} +import scala.collection.mutable + +private object ExecutableCallable { + + def getParameters(executable: Executable): Seq[JavaBeanParameter] = { + if (executable.getParameterCount == 0) Seq.empty[JavaBeanParameter] + else { + val parameters = new mutable.ArrayBuffer[JavaBeanParameter](executable.getParameterCount) + + val parameterArray = executable.getParameters + val parameterTypes = executable.getParameterTypes + val genericParameterTypes = executable.getGenericParameterTypes + if (parameterTypes.length == genericParameterTypes.length) { + var index = 0; + val length = parameterArray.length + while (index < length) { + parameters.append( + new JavaBeanParameter( + index, + parameterArray(index), + parameterTypes(index), + getErasedTypeIfTypeVariable(genericParameterTypes(index)))) + + index += 1 + } + } else { + // in this case, we have synthetic or implicit parameters + val hasParameterModifierInfo = isAnyParameterCarryingMetadata(parameterArray) + if (!hasParameterModifierInfo) + throw new ValidationException("") + + var explicitlyDeclaredParameterIndex = 0 + var index = 0; + val length = parameterArray.length + while (index < length) { + if (explicitlyDeclaredParameterIndex < genericParameterTypes.length + && isExplicit(parameterArray(index)) + && parameterTypesMatch( + parameterTypes(index), + genericParameterTypes(explicitlyDeclaredParameterIndex))) { + // in this case we have a parameter that is present and matches ("most likely") to the one in the generic parameter types list + parameters.append( + new JavaBeanParameter( + index, + parameterArray(index), + parameterTypes(index), + getErasedTypeIfTypeVariable(genericParameterTypes(explicitlyDeclaredParameterIndex)) + ) + ) + explicitlyDeclaredParameterIndex += 1 + } else { + // in this case, the parameter is not present in genericParameterTypes, or the types doesn't match + parameters.append( + new JavaBeanParameter( + index, + parameterArray(index), + parameterTypes(index), + parameterTypes(index))) + } + + index += 1 + } + } + + parameters.toSeq + } + } + + private[this] def getErasedTypeIfTypeVariable(genericType: Type): Type = genericType match { + case _: TypeVariable[_] => + TypeHelper.getErasedType(genericType) + case _ => + genericType + } + + private[this] def isAnyParameterCarryingMetadata(parameterArray: Array[Parameter]): Boolean = { + parameterArray.exists(p => p.isSynthetic || p.isImplicit) + } + + private[this] def parameterTypesMatch(paramType: Class[_], genericParamType: Type): Boolean = { + TypeHelper.getErasedType(genericParamType).equals(paramType) + } + + private[this] def isExplicit(parameter: Parameter): Boolean = { + !parameter.isSynthetic && !parameter.isImplicit + } +} + +abstract class ExecutableCallable[T <: Executable]( + val executable: T, + hasReturnValue: Boolean) + extends Callable { + + private val `type`: Type = ReflectionHelper.typeOf(executable) + private val parameters: Seq[JavaBeanParameter] = ExecutableCallable.getParameters(executable) + private val typeForValidatorResolution: Type = ReflectionHelper.boxedType(this.`type`) + + def getParameters: Seq[JavaBeanParameter] = this.parameters + override def getName: String = executable.getName + override def getDeclaringClass: Class[_] = executable.getDeclaringClass + override def getTypeForValidatorResolution: Type = this.typeForValidatorResolution + override def getType: Type = this.`type` + override def hasReturnValue: Boolean = hasReturnValue + override def hasParameters: Boolean = this.parameters.nonEmpty + override def getParameterCount: Int = this.parameters.size + override def getParameterGenericType(index: Int): Type = this.parameters(index).getGenericType + override def getParameterTypes: Array[Class[_]] = executable.getParameterTypes + override def getParameterName( + parameterNameProvider: ExecutableParameterNameProvider, + parameterIndex: Int + ): String = parameterNameProvider.getParameterNames(executable).get(parameterIndex) + + override def isPrivate: Boolean = Modifier.isPrivate(executable.getModifiers) + override def getSignature: Signature = { + val simpleName = executable match { + case c: Constructor[_] => + Classes.simpleName(c.getDeclaringClass) + case _ => executable.getName + } + ExecutableHelper.getSignature(simpleName, executable.getParameterTypes) + } + + override def overrides( + executableHelper: ExecutableHelper, + superTypeMethod: Callable + ): Boolean = executableHelper.overrides( + executable.asInstanceOf[Method], + superTypeMethod.asInstanceOf[ExecutableCallable[T]].executable.asInstanceOf[Method] + ) + + override def isResolvedToSameMethodInHierarchy( + executableHelper: ExecutableHelper, + mainSubType: Class[_], + superTypeMethod: Callable + ): Boolean = + executableHelper + .isResolvedToSameMethodInHierarchy( + mainSubType, + executable.asInstanceOf[Method], + superTypeMethod.asInstanceOf[ExecutableCallable[T]].executable.asInstanceOf[Method] + ) + + override def equals(o: Any): Boolean = { + if (this == o) { + true + } else if (o != null && this.getClass == o.getClass) { + val that: ExecutableCallable[T] = o.asInstanceOf[ExecutableCallable[T]] + if (this.hasReturnValue != that.hasReturnValue) { + false + } else if (!this.executable.equals(that.executable)) { + false + } else { + if (this.typeForValidatorResolution != that.typeForValidatorResolution) false + else this.`type` == that.`type` + } + } else { + false + } + } + + override def hashCode: Int = { + var result = executable.hashCode + result = 31 * result + this.typeForValidatorResolution.hashCode + result = 31 * result + (if (this.hasReturnValue) 1 else 0) + result = 31 * result + this.`type`.hashCode + result + } + + override def toString: String = + ExecutableHelper.getExecutableAsString( + this.getDeclaringClass.getSimpleName + "#" + executable.getName, + executable.getParameterTypes: _* + ) +} diff --git a/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/MethodCallable.scala b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/MethodCallable.scala new file mode 100644 index 0000000000..f096a527d5 --- /dev/null +++ b/util-validator/src/main/scala/org/hibernate/validator/internal/properties/javabean/MethodCallable.scala @@ -0,0 +1,13 @@ +package org.hibernate.validator.internal.properties.javabean + +import java.lang.reflect.{Method, TypeVariable} +import org.hibernate.validator.internal.metadata.raw.ConstrainedElement.ConstrainedElementKind + +class MethodCallable(method: Method) + extends ExecutableCallable[Method](method, method.getGenericReturnType != classOf[Unit]) { + + def getConstrainedElementKind: ConstrainedElementKind = ConstrainedElementKind.METHOD + + def getTypeParameters: Array[TypeVariable[Class[_]]] = + method.getReturnType.getTypeParameters.asInstanceOf[Array[TypeVariable[Class[_]]]] +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/BUILD.bazel b/util-validator/src/test/java/com/twitter/util/validation/BUILD.bazel new file mode 100644 index 0000000000..d6a155c79a --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/BUILD.bazel @@ -0,0 +1,16 @@ +java_library( + sources = [ + "*.java", + "constraints/*.java", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], + exports = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + ], +) diff --git a/util-validator/src/test/java/com/twitter/util/validation/CarMake.java b/util-validator/src/test/java/com/twitter/util/validation/CarMake.java new file mode 100644 index 0000000000..ef1fdd3468 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/CarMake.java @@ -0,0 +1,8 @@ +package com.twitter.util.validation; + +/** + * Test Java Enum + */ +public enum CarMake { + Ford, Audi, Volkswagen, Tesla, Honda; +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/DoesNothing.java b/util-validator/src/test/java/com/twitter/util/validation/DoesNothing.java new file mode 100644 index 0000000000..6b74435c6b --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/DoesNothing.java @@ -0,0 +1,16 @@ +package com.twitter.util.validation; + +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +public @interface DoesNothing {} diff --git a/util-validator/src/test/java/com/twitter/util/validation/constraints/InvalidConstraint.java b/util-validator/src/test/java/com/twitter/util/validation/constraints/InvalidConstraint.java new file mode 100644 index 0000000000..23b0926846 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/constraints/InvalidConstraint.java @@ -0,0 +1,29 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +public @interface InvalidConstraint { + + String message() default "invalid"; + + Class[] groups() default {}; + + Class[] payload() default {}; +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCount.java b/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCount.java new file mode 100644 index 0000000000..34b18f8419 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCount.java @@ -0,0 +1,46 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Repeatable; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +@Repeatable(ValidPassengerCount.List.class) +public @interface ValidPassengerCount { + + String message() default "number of passenger(s) is not valid"; + + Class[] groups() default {}; + + Class[] payload() default {}; + + /** + * The max number of passengers. + */ + long max(); + + /** + * Defines several {@code @OneOf} annotations on the same element. + */ + @Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) + @Retention(RUNTIME) + @Documented + public @interface List { + ValidPassengerCount[] value(); + } +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCountReturnValue.java b/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCountReturnValue.java new file mode 100644 index 0000000000..4065b33ef1 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/constraints/ValidPassengerCountReturnValue.java @@ -0,0 +1,46 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Repeatable; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.CONSTRUCTOR; +import static java.lang.annotation.ElementType.FIELD; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.ElementType.PARAMETER; +import static java.lang.annotation.ElementType.TYPE_USE; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Documented +@Constraint(validatedBy = {}) +@Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) +@Retention(RUNTIME) +@Repeatable(ValidPassengerCountReturnValue.List.class) +public @interface ValidPassengerCountReturnValue { + + String message() default "number of passenger(s) is not valid"; + + Class[] groups() default {}; + + Class[] payload() default {}; + + /** + * The max number of passengers. + */ + long max(); + + /** + * Defines several {@code @OneOf} annotations on the same element. + */ + @Target({METHOD, FIELD, ANNOTATION_TYPE, CONSTRUCTOR, PARAMETER, TYPE_USE}) + @Retention(RUNTIME) + @Documented + public @interface List { + ValidPassengerCountReturnValue[] value(); + } +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/BUILD.bazel b/util-validator/src/test/java/com/twitter/util/validation/tests/BUILD.bazel new file mode 100644 index 0000000000..85128f9f56 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/BUILD.bazel @@ -0,0 +1,27 @@ +junit_tests( + sources = ["**/*.java"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/junit", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/slf4j:slf4j-api", + scoped( + "3rdparty/jvm/org/slf4j:slf4j-simple", + scope = "runtime", + ), + "util/util-slf4j-api/src/main/scala", + "util/util-validator/src/main/scala/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation/cfg", + "util/util-validator/src/main/scala/com/twitter/util/validation/conversions", + "util/util-validator/src/main/scala/com/twitter/util/validation/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/executable", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/validators", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + "util/util-validator/src/test/java/com/twitter/util/validation", + ], +) diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/Car.java b/util-validator/src/test/java/com/twitter/util/validation/tests/Car.java new file mode 100644 index 0000000000..26543766e9 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/Car.java @@ -0,0 +1,52 @@ +package com.twitter.util.validation.tests; + +import java.util.Collections; +import java.util.List; + +import com.twitter.util.validation.constraints.ValidPassengerCount; + +import jakarta.validation.constraints.Min; +import jakarta.validation.constraints.NotNull; +import jakarta.validation.constraints.Size; + +@ValidPassengerCount(max = 2) +public class Car { + @NotNull + private final String manufacturer; + + @NotNull + @Size(min = 2, max = 14) + private final String licensePlate; + + @Min(2) + private final int seatCount; + + private List passengers; + + public Car(String manufacturer, String licencePlate, int seatCount) { + this.manufacturer = manufacturer; + this.licensePlate = licencePlate; + this.seatCount = seatCount; + this.passengers = Collections.emptyList(); + } + + public void load(List passengers) { + this.passengers = passengers; + } + + public String getManufacturer() { + return manufacturer; + } + + public String getLicensePlate() { + return licensePlate; + } + + public int getSeatCount() { + return seatCount; + } + + public List getPassengers() { + return passengers; + } +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/Person.java b/util-validator/src/test/java/com/twitter/util/validation/tests/Person.java new file mode 100644 index 0000000000..bfc5872813 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/Person.java @@ -0,0 +1,16 @@ +package com.twitter.util.validation.tests; + +import jakarta.validation.constraints.NotNull; + +public class Person { + @NotNull + private String name; + + public Person(String name) { + this.name = name; + } + + public String getName() { + return name; + } +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/ScalaValidatorJavaTest.java b/util-validator/src/test/java/com/twitter/util/validation/tests/ScalaValidatorJavaTest.java new file mode 100644 index 0000000000..3fa03564a9 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/ScalaValidatorJavaTest.java @@ -0,0 +1,155 @@ +package com.twitter.util.validation.tests; + +import java.lang.annotation.Annotation; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.HashMap; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; + +import org.junit.Assert; +import org.junit.Test; + +import com.twitter.util.validation.ScalaValidator; +import com.twitter.util.validation.cfg.ConstraintMapping; +import com.twitter.util.validation.constraints.ValidPassengerCount; + +import jakarta.validation.ConstraintViolation; + +public class ScalaValidatorJavaTest extends Assert { + private final ConstraintMapping mapping = + ConstraintMapping.apply( + ValidPassengerCount.class, + ValidPassengerCountConstraintValidator.class, false + ); + + private final ScalaValidator validator = + ScalaValidator.builder() + .withConstraintMapping(mapping) + .validator(); + + @Test + public void testConstructors() { + Assert.assertNotNull(ScalaValidator.apply()); + Assert.assertNotNull(ScalaValidator.builder().validator()); + } + + @Test + public void testBuilder() { + Assert.assertNotNull(ScalaValidator.builder().validator()); + + Set mappings = new HashSet<>(); + mappings.add(mapping); + + Assert.assertNotNull( + ScalaValidator.builder() + .withConstraintMapping(mapping) + .validator() + ); + + Assert.assertNotNull( + ScalaValidator.builder() + .withConstraintMappings(mappings) + .validator() + ); + } + + @Test + public void testUnderlyingValidator() { + Car car = new Car("Renault", "DD-AB-123", 4); + + Set> violations = validator.underlying().validate(car); + Assert.assertEquals(violations.size(), 0); + + List passengers = new ArrayList(); + passengers.add(new Person("R Smith")); + passengers.add(new Person("S Smith")); + passengers.add(new Person("T Smith")); + passengers.add(new Person("U Smith")); + + car = new Car("Renault", "DD-AB-123", 2); + car.load(passengers); + + violations = validator.underlying().validate(car); + Assert.assertEquals(violations.size(), 1); + } + + @Test + public void testValidateFieldValue() { + Map sizeConstraintValues = new HashMap<>(); + sizeConstraintValues.put("min", 5); + sizeConstraintValues.put("max", 7); + + Map, Map> constraints = new HashMap<>(); + constraints.put(jakarta.validation.constraints.Size.class, sizeConstraintValues); + + Set> violations = + validator.validateFieldValue( + constraints, + "javaSetIntField", + newSet(6) + ); + + Assert.assertEquals(violations.size(), 0); + + Set value = newSet(9); + violations = + validator.validateFieldValue( + constraints, + "javaSetIntField", + value + ); + + Assert.assertEquals(violations.size(), 1); + ConstraintViolation violation = violations.iterator().next(); + Assert.assertEquals(violation.getPropertyPath().toString(), "javaSetIntField"); + Assert.assertEquals(violation.getMessage(), "size must be between 5 and 7"); + Assert.assertEquals(violation.getInvalidValue(), value); + } + + @Test + public void testValidateFieldValueWithGroups() { + Class[] groups = {TestConstraintGroup.class}; + Map sizeConstraintValues = new HashMap<>(); + sizeConstraintValues.put("min", 5); + sizeConstraintValues.put("max", 7); + sizeConstraintValues.put("groups", groups); + + Map, Map> constraints = new HashMap<>(); + constraints.put(jakarta.validation.constraints.Size.class, sizeConstraintValues); + + Set value = newSet(9); + // TestConstraintGroup is not active + Set> violations = + validator.validateFieldValue( + constraints, + "javaSetIntField", + value + ); + Assert.assertEquals(violations.size(), 0); + + // TestConstraintGroup is active + violations = + validator.validateFieldValue( + constraints, + "javaSetIntField", + value, + TestConstraintGroup.class + ); + ConstraintViolation violation = violations.iterator().next(); + Assert.assertEquals(violation.getPropertyPath().toString(), "javaSetIntField"); + Assert.assertEquals(violation.getMessage(), "size must be between 5 and 7"); + Assert.assertEquals(violation.getInvalidValue(), value); + } + + public Set newSet(Integer total) { + Set results = new HashSet<>(); + for (int i = 0; i < total; i++) { + results.add(i); + } + + return results; + } +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/TestConstraintGroup.java b/util-validator/src/test/java/com/twitter/util/validation/tests/TestConstraintGroup.java new file mode 100644 index 0000000000..100deedc95 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/TestConstraintGroup.java @@ -0,0 +1,4 @@ +package com.twitter.util.validation.tests; + +public interface TestConstraintGroup { +} diff --git a/util-validator/src/test/java/com/twitter/util/validation/tests/ValidPassengerCountConstraintValidator.java b/util-validator/src/test/java/com/twitter/util/validation/tests/ValidPassengerCountConstraintValidator.java new file mode 100644 index 0000000000..4e4fb5f062 --- /dev/null +++ b/util-validator/src/test/java/com/twitter/util/validation/tests/ValidPassengerCountConstraintValidator.java @@ -0,0 +1,25 @@ +package com.twitter.util.validation.tests; + +import com.twitter.util.validation.constraints.ValidPassengerCount; + +import jakarta.validation.ConstraintValidator; +import jakarta.validation.ConstraintValidatorContext; + +public class ValidPassengerCountConstraintValidator implements ConstraintValidator { + + private long max; + + @Override + public void initialize(ValidPassengerCount constraintAnnotation) { + this.max = constraintAnnotation.max(); + } + + @Override + public boolean isValid(Car car, ConstraintValidatorContext context) { + if ( car == null ) { + return true; + } + + return car.getPassengers().size() <= this.max; + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/AssertViolationTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/AssertViolationTest.scala new file mode 100644 index 0000000000..9061806fc9 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/AssertViolationTest.scala @@ -0,0 +1,59 @@ +package com.twitter.util.validation + +import com.twitter.util.validation.engine.ConstraintViolationHelper +import jakarta.validation.ConstraintViolation +import java.lang.annotation.Annotation +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import scala.reflect.runtime.universe._ + +object AssertViolationTest { + case class WithViolation[T]( + path: String, + message: String, + invalidValue: Any, + rootBeanClazz: Class[T], + root: T, + annotation: Option[Annotation] = None) +} + +@RunWith(classOf[JUnitRunner]) +abstract class AssertViolationTest extends AnyFunSuite with Matchers { + import AssertViolationTest._ + + protected def validator: ScalaValidator + + protected def assertViolations[T: TypeTag]( + obj: T, + groups: Seq[Class[_]] = Seq.empty, + withViolations: Seq[WithViolation[T]] = Seq.empty + ): Unit = { + val violations: Set[ConstraintViolation[T]] = validator.validate(obj, groups: _*) + assertViolations(violations, withViolations) + } + + protected def assertViolations[T: TypeTag]( + violations: Set[ConstraintViolation[T]], + withViolations: Seq[WithViolation[T]] + ): Unit = { + violations.size should equal(withViolations.size) + val sortedViolations: Seq[ConstraintViolation[T]] = + ConstraintViolationHelper.sortViolations(violations) + for ((constraintViolation, index) <- sortedViolations.zipWithIndex) { + val withViolation = withViolations(index) + constraintViolation.getMessage should equal(withViolation.message) + if (constraintViolation.getPropertyPath == null) withViolation.path should be(null) + else constraintViolation.getPropertyPath.toString should equal(withViolation.path) + constraintViolation.getInvalidValue should equal(withViolation.invalidValue) + constraintViolation.getRootBeanClass should equal(withViolation.rootBeanClazz) + withViolation.root.getClass.getName == constraintViolation.getRootBean + .asInstanceOf[T].getClass.getName should be(true) + constraintViolation.getLeafBean should not be null + withViolation.annotation.foreach { ann => + constraintViolation.getConstraintDescriptor.getAnnotation should equal(ann) + } + } + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/BUILD.bazel b/util-validator/src/test/scala/com/twitter/util/validation/BUILD.bazel new file mode 100644 index 0000000000..c9a121bb32 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/BUILD.bazel @@ -0,0 +1,37 @@ +junit_tests( + sources = [ + "**/*.java", + "**/*.scala", + ], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/com/fasterxml/jackson/core:jackson-annotations", + "3rdparty/jvm/com/github/ben-manes/caffeine", + "3rdparty/jvm/jakarta/validation:jakarta.validation-api", + "3rdparty/jvm/org/hibernate/validator:hibernate-validator", + "3rdparty/jvm/org/json4s:json4s-core", + "3rdparty/jvm/org/scalacheck", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:scalacheck", + scoped( + "3rdparty/jvm/org/slf4j:slf4j-simple", + scope = "runtime", + ), + "util/util-core:scala", + "util/util-reflect/src/main/scala/com/twitter/util/reflect", + "util/util-slf4j-api/src/main/scala", + "util/util-validator/src/main/scala/com/twitter/util/validation", + "util/util-validator/src/main/scala/com/twitter/util/validation/cfg", + "util/util-validator/src/main/scala/com/twitter/util/validation/conversions", + "util/util-validator/src/main/scala/com/twitter/util/validation/engine", + "util/util-validator/src/main/scala/com/twitter/util/validation/executable", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal", + "util/util-validator/src/main/scala/com/twitter/util/validation/internal/validators", + "util/util-validator/src/main/scala/com/twitter/util/validation/metadata", + "util/util-validator/src/test/java/com/twitter/util/validation", + ], +) diff --git a/util-validator/src/test/scala/com/twitter/util/validation/OptionUnwrappingTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/OptionUnwrappingTest.scala new file mode 100644 index 0000000000..634d346a25 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/OptionUnwrappingTest.scala @@ -0,0 +1,170 @@ +package com.twitter.util.validation + +import com.twitter.util.validation.AssertViolationTest._ +import com.twitter.util.validation.caseclasses._ + +/** + * Options are "unwrapped" by default for validation, meaning any supplied constraint + * annotation is applied to the contained type/value of the Option. The unwrapping + * happens through the Validator and not in the individual ConstraintValidator thus we + * test though a Validator here and why these tests are not in the individual ConstraintValidator + * tests. + */ +class OptionUnwrappingTest extends AssertViolationTest { + + override protected val validator: ScalaValidator = ScalaValidator.builder.validator + + test("CountryCode") { + assertViolations( + obj = CountryCodeOptionExample(Some("11")), + withViolations = Seq( + WithViolation( + "countryCode", + "11 not a valid country code", + "11", + classOf[CountryCodeOptionExample], + CountryCodeOptionExample(Some("11"))) + ) + ) + assertViolations(obj = CountryCodeOptionExample(None)) + + assertViolations( + obj = CountryCodeOptionSeqExample(Some(Seq("11", "22"))), + withViolations = Seq( + WithViolation( + "countryCode", + "[11, 22] not a valid country code", + Seq("11", "22"), + classOf[CountryCodeOptionSeqExample], + CountryCodeOptionSeqExample(Some(Seq("11", "22"))) + ) + ) + ) + assertViolations(obj = CountryCodeOptionSeqExample(None)) + + assertViolations( + obj = CountryCodeOptionArrayExample(Some(Array("11", "22"))), + withViolations = Seq( + WithViolation( + "countryCode", + "[11, 22] not a valid country code", + Array("11", "22"), + classOf[CountryCodeOptionArrayExample], + CountryCodeOptionArrayExample(Some(Array("11", "22"))) + ) + ) + ) + assertViolations(obj = CountryCodeOptionArrayExample(None)) + + assertViolations( + obj = CountryCodeOptionInvalidTypeExample(Some(12345L)), + withViolations = Seq( + WithViolation( + "countryCode", + "12345 not a valid country code", + 12345L, + classOf[CountryCodeOptionInvalidTypeExample], + CountryCodeOptionInvalidTypeExample(Some(12345L))) + ) + ) + assertViolations(obj = CountryCodeOptionInvalidTypeExample(None)) + } + + test("OneOf") { + assertViolations( + obj = OneOfOptionExample(Some("g")), + withViolations = Seq( + WithViolation( + "enumValue", + "g not one of [a, B, c]", + "g", + classOf[OneOfOptionExample], + OneOfOptionExample(Some("g"))) + ) + ) + assertViolations(OneOfOptionExample(None)) + + assertViolations( + obj = OneOfOptionSeqExample(Some(Seq("g", "h", "i"))), + withViolations = Seq( + WithViolation( + "enumValue", + "[g, h, i] not one of [a, B, c]", + Seq("g", "h", "i"), + classOf[OneOfOptionSeqExample], + OneOfOptionSeqExample(Some(Seq("g", "h", "i")))) + ) + ) + assertViolations(OneOfOptionSeqExample(None)) + + assertViolations( + obj = OneOfOptionInvalidTypeExample(Some(10L)), + withViolations = Seq( + WithViolation( + "enumValue", + "10 not one of [a, B, c]", + 10L, + classOf[OneOfOptionInvalidTypeExample], + OneOfOptionInvalidTypeExample(Some(10L))) + ) + ) + assertViolations(OneOfOptionSeqExample(None)) + assertViolations(OneOfOptionInvalidTypeExample(None)) + } + + test("UUID") { + assertViolations( + obj = UUIDOptionExample(Some("abcd1234")), + withViolations = Seq( + WithViolation( + "uuid", + "must be a valid UUID", + "abcd1234", + classOf[UUIDOptionExample], + UUIDOptionExample(Some("abcd1234"))) + ) + ) + assertViolations(obj = UUIDOptionExample(None)) + } + + test("NestedOptionExample") { + assertViolations( + obj = NestedOptionExample( + "abcd1234", + "Foo Var", + None + ), + ) + + val value = NestedOptionExample( + "", + "Foo Var", + Some( + User( + "", + "Bar Var", + "H" + ) + ) + ) + assertViolations( + obj = value, + withViolations = Seq( + WithViolation( + "id", + "must not be empty", + "", + classOf[NestedOptionExample], + value + ), + WithViolation( + "user.gender", + "H not one of [F, M, Other]", + "H", + classOf[NestedOptionExample], + value), + WithViolation("user.id", "must not be empty", "", classOf[NestedOptionExample], value) + ) + ) + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/ScalaValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/ScalaValidatorTest.scala new file mode 100644 index 0000000000..71609b96c1 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/ScalaValidatorTest.scala @@ -0,0 +1,1528 @@ +package com.twitter.util.validation + +import com.twitter.util.validation.AssertViolationTest.WithViolation +import com.twitter.util.validation.caseclasses._ +import com.twitter.util.validation.cfg.ConstraintMapping +import com.twitter.util.validation.constraints.{ + CountryCode, + InvalidConstraint, + StateConstraint, + StateConstraintPayload, + ValidPassengerCount +} +import com.twitter.util.validation.conversions.ConstraintViolationOps._ +import com.twitter.util.validation.internal.AnnotationFactory +import com.twitter.util.validation.metadata.{CaseClassDescriptor, PropertyDescriptor} +import com.twitter.util.validation.internal.validators.{ + ISO3166CountryCodeConstraintValidator, + InvalidConstraintValidator, + NotEmptyAnyConstraintValidator, + NotEmptyPathConstraintValidator, + ValidPassengerCountConstraintValidator +} +import jakarta.validation.constraints.{Min, NotEmpty, Pattern, Size} +import jakarta.validation.{ + ConstraintDeclarationException, + ConstraintValidator, + ConstraintViolation, + ConstraintViolationException, + UnexpectedTypeException, + Validation, + ValidationException +} +import java.lang.annotation.Annotation +import java.time.LocalDate +import java.util +import java.util.{UUID => JUUID} +import org.hibernate.validator.HibernateValidatorConfiguration +import org.json4s.reflect.{ClassDescriptor, ConstructorDescriptor, Reflector} +import org.scalatest.BeforeAndAfterAll +import scala.jdk.CollectionConverters._ + +private object ScalaValidatorTest { + val DefaultAddress: Address = + Address(line1 = "1234 Main St", city = "Anywhere", state = "CA", zipcode = "94102") + + val CustomConstraintMappings: Set[ConstraintMapping] = + Set( + ConstraintMapping( + classOf[ValidPassengerCount], + classOf[ValidPassengerCountConstraintValidator], + includeExistingValidators = false + ), + ConstraintMapping( + classOf[InvalidConstraint], + classOf[InvalidConstraintValidator] + ) + ) +} + +class ScalaValidatorTest extends AssertViolationTest with BeforeAndAfterAll { + import ScalaValidatorTest._ + + override protected val validator: ScalaValidator = + ScalaValidator.builder + .withConstraintMappings(CustomConstraintMappings) + .validator + + override protected def afterAll(): Unit = { + validator.close() + super.afterAll() + } + + test("ScalaValidator#with custom message interpolator") { + val withCustomMessageInterpolator = ScalaValidator.builder + .withMessageInterpolator(new WrongMessageInterpolator) + .validator + try { + val car = SmallCar(null, "DD-AB-123", 4) + val violations = withCustomMessageInterpolator.validate(car) + assertViolations( + violations = violations, + withViolations = Seq( + WithViolation( + "manufacturer", + "Whatever you entered was wrong", + null, + classOf[SmallCar], + car) + ) + ) + } finally { + withCustomMessageInterpolator.close() + } + } + + test("ScalaValidator#nested case class definition") { + val value = OuterObject.InnerObject.SuperNestedCaseClass(4L) + + val violations = validator.validate(value) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should be("id") + violation.getMessage should be("must be greater than or equal to 5") + violation.getRootBeanClass should equal(classOf[OuterObject.InnerObject.SuperNestedCaseClass]) + violation.getRootBean should equal(value) + violation.getInvalidValue should equal(4L) + } + + test("ScalaValidator#unicode field name") { + val value = UnicodeNameCaseClass(5, "") + + assertViolations( + obj = value, + withViolations = Seq( + WithViolation( + "name", + "must not be empty", + "", + classOf[UnicodeNameCaseClass], + value, + getValidationAnnotations(classOf[UnicodeNameCaseClass], "name").lastOption + ), + WithViolation( + "winning-id", + "must be less than or equal to 3", + 5, + classOf[UnicodeNameCaseClass], + value, + getValidationAnnotations(classOf[UnicodeNameCaseClass], "winning-id").lastOption + ), + ) + ) + } + + test("ScalaValidator#correctly handle not cascading for a non-case class type") { + val value = WithBogusCascadedField("DD-AB-123", Seq(1, 2, 3)) + assertViolations(value) + } + + test("ScalaValidator#with InvalidConstraint") { + val value = AlwaysFails("fails") + + val e = intercept[RuntimeException] { + validator.validate(value) + } + e.getMessage should equal("validator foo error") + } + + test("ScalaValidator#with failing MethodValidation") { + // fails the @NotEmpty on the id field but the method validation blows up + val value = AlwaysFailsMethodValidation("") + + val e = intercept[ValidationException] { + validator.validate(value) + } + e.getMessage should equal("oh noes!") + } + + test("ScalaValidator#isConstraintValidator") { + val stateConstraint = + AnnotationFactory.newInstance(classOf[StateConstraint], Map.empty) + validator.isConstraintAnnotation(stateConstraint) should be(true) + } + + test("ScalaValidator#constraint mapping 1") { + val value = PathNotEmpty(Path.Empty, "abcd1234") + + // no validator registered `@NotEmpty` for Path type + intercept[UnexpectedTypeException] { + validator.validate(value) + } + + // define an additional constraint validator for `@NotEmpty` which works for + // Path types. note we include all existing validators + val pathCustomConstraintMapping = + ConstraintMapping(classOf[NotEmpty], classOf[NotEmptyPathConstraintValidator]) + + val withPathConstraintValidator: ScalaValidator = + ScalaValidator.builder + .withConstraintMapping(pathCustomConstraintMapping).validator + + try { + val violations = withPathConstraintValidator.validate(value) + violations.size should equal(1) + violations.head.getPropertyPath.toString should be("path") + violations.head.getMessage should be("must not be empty") + violations.head.getRootBeanClass should equal(classOf[PathNotEmpty]) + violations.head.getRootBean should equal(value) + violations.head.getInvalidValue should equal(Path.Empty) + + // valid instance + withPathConstraintValidator + .validate(PathNotEmpty(Path("node"), "abcd1234")) + .isEmpty should be(true) + } finally { + withPathConstraintValidator.close() + } + + // define an additional constraint validator for `@NotEmpty` which works for + // Path types. note we do include all existing validators which will cause validation to + // fail since we can no longer find a validator for `@NotEmpty` which works for string + val pathCustomConstraintMapping2 = + ConstraintMapping( + classOf[NotEmpty], + classOf[NotEmptyPathConstraintValidator], + includeExistingValidators = false) + + val withPathConstraintValidator2: ScalaValidator = + ScalaValidator.builder + .withConstraintMapping(pathCustomConstraintMapping2).validator + try { + intercept[UnexpectedTypeException] { + withPathConstraintValidator2.validate(value) + } + } finally { + withPathConstraintValidator2.close() + } + } + + test("ScalaValidator#constraint mapping 2") { + // attempting to add a mapping for the same annotation multiple times will result in an error + val pathCustomConstraintMapping1 = + ConstraintMapping(classOf[NotEmpty], classOf[NotEmptyPathConstraintValidator]) + val pathCustomConstraintMapping2 = + ConstraintMapping(classOf[NotEmpty], classOf[NotEmptyAnyConstraintValidator]) + + val e1 = intercept[ValidationException] { + ScalaValidator.builder + .withConstraintMappings(Set(pathCustomConstraintMapping1, pathCustomConstraintMapping2)) + .validator + } + e1.getMessage.endsWith( + "jakarta.validation.constraints.NotEmpty is configured more than once via the programmatic constraint definition API.") should be( + true) + + val pathCustomConstraintMapping3 = + ConstraintMapping(classOf[NotEmpty], classOf[NotEmptyPathConstraintValidator]) + val pathCustomConstraintMapping4 = + ConstraintMapping( + classOf[NotEmpty], + classOf[NotEmptyAnyConstraintValidator], + includeExistingValidators = false) + + val e2 = intercept[ValidationException] { + ScalaValidator.builder + .withConstraintMappings(Set(pathCustomConstraintMapping3, pathCustomConstraintMapping4)) + .validator + } + + e2.getMessage.endsWith( + "jakarta.validation.constraints.NotEmpty is configured more than once via the programmatic constraint definition API.") should be( + true) + + val pathCustomConstraintMapping5 = + ConstraintMapping( + classOf[NotEmpty], + classOf[NotEmptyPathConstraintValidator], + includeExistingValidators = false) + val pathCustomConstraintMapping6 = + ConstraintMapping(classOf[NotEmpty], classOf[NotEmptyAnyConstraintValidator]) + + val e3 = intercept[ValidationException] { + ScalaValidator.builder + .withConstraintMappings(Set(pathCustomConstraintMapping5, pathCustomConstraintMapping6)) + .validator + } + + e3.getMessage.endsWith( + "jakarta.validation.constraints.NotEmpty is configured more than once via the programmatic constraint definition API.") should be( + true) + + val pathCustomConstraintMapping7 = + ConstraintMapping( + classOf[NotEmpty], + classOf[NotEmptyPathConstraintValidator], + includeExistingValidators = false) + val pathCustomConstraintMapping8 = + ConstraintMapping( + classOf[NotEmpty], + classOf[NotEmptyAnyConstraintValidator], + includeExistingValidators = false) + + val e4 = intercept[ValidationException] { + ScalaValidator.builder + .withConstraintMappings(Set(pathCustomConstraintMapping7, pathCustomConstraintMapping8)) + .validator + } + + e4.getMessage.endsWith( + "jakarta.validation.constraints.NotEmpty is configured more than once via the programmatic constraint definition API.") should be( + true) + } + + test("ScalaValidator#payload") { + val address = + Address(line1 = "1234 Main St", city = "Anywhere", state = "PA", zipcode = "94102") + val violations = validator.validate(address) + violations.size should equal(1) + val violation = violations.head + val payload = violation.getDynamicPayload(classOf[StateConstraintPayload]).orNull + payload == null should be(false) + payload.isInstanceOf[StateConstraintPayload] should be(true) + payload should equal(StateConstraintPayload("PA", "CA")) + } + + test("ScalaValidator#method validation payload") { + val user = NestedUser( + id = "abcd1234", + person = Person("abcd1234", "R. Franklin", DefaultAddress), + gender = "F", + job = "" + ) + val violations = validator.validate(user) + violations.size should equal(1) + val violation = violations.head + val payload = violation.getDynamicPayload(classOf[NestedUserPayload]).orNull + payload == null should be(false) + payload.isInstanceOf[NestedUserPayload] should be(true) + payload should equal(NestedUserPayload("abcd1234", "")) + } + + test("ScalaValidator#newInstance") { + val stateConstraint = + AnnotationFactory.newInstance(classOf[StateConstraint], Map.empty) + validator.isConstraintAnnotation(stateConstraint) should be(true) + + val minConstraint = AnnotationFactory.newInstance( + classOf[Min], + Map("value" -> 1L) + ) + minConstraint.groups().isEmpty should be(true) + minConstraint.message() should equal("{jakarta.validation.constraints.Min.message}") + minConstraint.payload().isEmpty should be(true) + minConstraint.value() should equal(1L) + validator.isConstraintAnnotation(minConstraint) should be(true) + + val patternConstraint = AnnotationFactory.newInstance( + classOf[Pattern], + Map("regexp" -> ".*", "flags" -> Array[Pattern.Flag](Pattern.Flag.CASE_INSENSITIVE)) + ) + patternConstraint.groups().isEmpty should be(true) + patternConstraint.message() should equal("{jakarta.validation.constraints.Pattern.message}") + patternConstraint.payload().isEmpty should be(true) + patternConstraint.regexp() should equal(".*") + patternConstraint.flags() should equal(Array[Pattern.Flag](Pattern.Flag.CASE_INSENSITIVE)) + validator.isConstraintAnnotation(patternConstraint) should be(true) + } + + test("ScalaValidator#constraint validators by annotation class") { + val validators = validator.findConstraintValidators(classOf[CountryCode]) + validators.size should equal(1) + validators.head.getClass.getName should be( + classOf[ISO3166CountryCodeConstraintValidator].getName) + + validators.head + .asInstanceOf[ConstraintValidator[CountryCode, String]].isValid("US", null) should be(true) + } + + test("ScalaValidator#nulls") { + val car = SmallCar(null, "DD-AB-123", 4) + assertViolations( + obj = car, + withViolations = Seq( + WithViolation("manufacturer", "must not be null", null, classOf[SmallCar], car) + ) + ) + } + + test("ScalaValidator#not a case class") { + val value = new NotACaseClass(null, "DD-AB-123", 4, "NotAUUID") + val e = intercept[ValidationException] { + validator.validate(value) + } + e.getMessage should equal(s"${classOf[NotACaseClass]} is not a valid case class.") + + // underlying validator doesn't find the Scala annotations (either executable nor inherited) + val violations = validator.underlying.validate(value).asScala.toSet + violations.isEmpty should be(true) + } + + test("ScalaValidator#works with meta annotation 1") { + val value = WithMetaFieldAnnotation("") + + assertViolations( + obj = value, + withViolations = Seq( + WithViolation( + "id", + "must not be empty", + "", + classOf[WithMetaFieldAnnotation], + value, + getValidationAnnotations(classOf[WithMetaFieldAnnotation], "id").headOption + ) + ) + ) + + // also works with underlying validator + val violations = validator.underlying.validate(value).asScala.toSet + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("id") + violation.getInvalidValue should equal("") + violation.getRootBeanClass should equal(classOf[WithMetaFieldAnnotation]) + violation.getRootBean should equal(value) + violation.getLeafBean should equal(value) + } + + test("ScalaValidator#works with meta annotation 2") { + val value = WithPartialMetaFieldAnnotation("") + + assertViolations( + obj = value, + withViolations = Seq( + WithViolation( + "id", + "must not be empty", + "", + classOf[WithPartialMetaFieldAnnotation], + value, + getValidationAnnotations(classOf[WithPartialMetaFieldAnnotation], "id").lastOption + ), + WithViolation( + "id", + "size must be between 2 and 10", + "", + classOf[WithPartialMetaFieldAnnotation], + value, + getValidationAnnotations(classOf[WithPartialMetaFieldAnnotation], "id").headOption + ) + ) + ) + } + + test("ScalaValidator#validateValue 1") { + val violations = validator.validateValue(classOf[OneOfSeqExample], "enumValue", Seq("g")) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("g not one of [a, B, c]") + violation.getPropertyPath.toString should equal("enumValue") + violation.getInvalidValue should equal(Seq("g")) + violation.getRootBeanClass should equal(classOf[OneOfSeqExample]) + violation.getRootBean should be(null) + violation.getLeafBean should be(null) + } + + test("ScalaValidator#validateValue 2") { + val descriptor = validator.getConstraintsForClass(classOf[OneOfSeqExample]) + val violations = validator.validateValue(descriptor, "enumValue", Seq("g")) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("g not one of [a, B, c]") + violation.getPropertyPath.toString should equal("enumValue") + violation.getInvalidValue should equal(Seq("g")) + violation.getRootBeanClass should equal(classOf[OneOfSeqExample]) + violation.getRootBean should be(null) + violation.getLeafBean should be(null) + } + + test("ScalaValidator#validateProperty 1") { + val value = CustomerAccount("", None) + val violations = validator.validateProperty(value, "accountName") + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("accountName") + violation.getInvalidValue should equal("") + violation.getRootBeanClass should equal(classOf[CustomerAccount]) + violation.getRootBean should equal(value) + violation.getLeafBean should equal(value) + } + + test("ScalaValidator#validateProperty 2") { + validator + .validateProperty( + MinIntExample(numberValue = 2), + "numberValue" + ).isEmpty should be(true) + assertViolations(obj = MinIntExample(numberValue = 2)) + + val invalid = MinIntExample(numberValue = 0) + val violations = validator + .validateProperty( + invalid, + "numberValue" + ) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should equal("numberValue") + violation.getMessage should equal("must be greater than or equal to 1") + violation.getRootBeanClass should equal(classOf[MinIntExample]) + violation.getRootBean.getClass.getName should equal(invalid.getClass.getName) + + assertViolations( + obj = invalid, + withViolations = Seq( + WithViolation( + "numberValue", + "must be greater than or equal to 1", + 0, + classOf[MinIntExample], + invalid, + getValidationAnnotations(classOf[MinIntExample], "numberValue").headOption + ) + ) + ) + } + + test("ScalaValidator#validateFieldValue") { + val constraints: Map[Class[_ <: Annotation], Map[String, Any]] = + Map(classOf[jakarta.validation.constraints.Size] -> Map("min" -> 5, "max" -> 7)) + + // size doesn't work for ints + intercept[UnexpectedTypeException] { + validator.validateFieldValue(constraints, "intField", 4) + } + + // should work fine for collections of the right size + validator + .validateFieldValue(constraints, "javaSetIntField", newJavaSet(6)).isEmpty should be(true) + validator + .validateFieldValue(constraints, "seqIntField", Seq(1, 2, 3, 4, 5, 6)).isEmpty should be(true) + validator + .validateFieldValue(constraints, "setIntField", Set(1, 2, 3, 4, 5, 6)).isEmpty should be(true) + validator + .validateFieldValue(constraints, "arrayIntField", Array(1, 2, 3, 4, 5, 6)).isEmpty should be( + true) + validator + .validateFieldValue(constraints, "stringField", "123456").isEmpty should be(true) + + // will fail for types of the wrong size + val violations = + validator.validateFieldValue(constraints, "seqIntField", Seq(1, 2, 3)) + violations.size should equal(1) + val violation: ConstraintViolation[Any] = violations.head + violation.getMessage should equal("size must be between 5 and 7") + violation.getPropertyPath.toString should equal("seqIntField") + violation.getInvalidValue should equal(Seq(1, 2, 3)) + violation.getRootBeanClass should equal(classOf[AnyRef]) + violation.getRootBean == null should be(true) + violation.getLeafBean should be(null) + + val invalidJavaSet = newJavaSet(3) + val violations2 = validator + .validateFieldValue(constraints, "javaSetIntField", invalidJavaSet) + violations2.size should equal(1) + val violation2: ConstraintViolation[Any] = violations.head + violation2.getMessage should equal("size must be between 5 and 7") + violation2.getPropertyPath.toString should equal("seqIntField") + violation2.getInvalidValue should equal(List(1, 2, 3)) + violation2.getRootBeanClass should equal(classOf[AnyRef]) + violation2.getRootBean == null should be(true) + violation2.getLeafBean should be(null) + + val constraintsWithGroup: Map[Class[_ <: Annotation], Map[String, Any]] = + Map( + classOf[jakarta.validation.constraints.Size] -> Map( + "min" -> 5, + "max" -> 7, + "groups" -> Array(classOf[PersonCheck]) // groups MUST be an array type + ) + ) + + // constraint group is not passed (not activated) + validator + .validateFieldValue(constraintsWithGroup, "seqIntField", Seq(1, 2, 3)).isEmpty should be(true) + // constraint group is passed (activated) + val violationsWithGroup = validator + .validateFieldValue(constraintsWithGroup, "seqIntField", Seq(1, 2, 3), classOf[PersonCheck]) + violationsWithGroup.size should equal(1) + val violationWg = violationsWithGroup.head + violationWg.getMessage should equal("size must be between 5 and 7") + violationWg.getPropertyPath.toString should equal("seqIntField") + violationWg.getInvalidValue should equal(Seq(1, 2, 3)) + violationWg.getRootBeanClass should equal(classOf[AnyRef]) + violationWg.getRootBean == null should be(true) + violationWg.getLeafBean should be(null) + + val violationsSt = + validator.validateFieldValue(constraints, "setIntField", Set(1, 2, 3, 4, 5, 6, 7, 8)) + violationsSt.size should equal(1) + val violationSt = violationsSt.head + violationSt.getMessage should equal("size must be between 5 and 7") + violationSt.getPropertyPath.toString should equal("setIntField") + violationSt.getInvalidValue should equal(Set(1, 2, 3, 4, 5, 6, 7, 8)) + violationSt.getRootBeanClass should equal(classOf[AnyRef]) + violationSt.getRootBean == null should be(true) + violationSt.getLeafBean should be(null) + + validator + .validateFieldValue( + Map[Class[_ <: Annotation], Map[String, AnyRef]](classOf[NotEmpty] + .asInstanceOf[Class[_ <: Annotation]] -> Map[String, AnyRef]()), + "data", + "Hello, world").isEmpty should be(true) + } + + test("ScalaValidator#groups support") { + // see: https://docs.jboss.org/hibernate/stable/scalaValidator/reference/en-US/html_single/#chapter-groups + val toTest = WithPersonCheck("", "Jane Doe") + // PersonCheck group not enabled, no violations + assertViolations(obj = toTest) + // Completely diff group enabled, no violations + validator.validate(obj = toTest, groups = classOf[OtherCheck]) + // PersonCheck group enabled + assertViolations( + obj = toTest, + groups = Seq(classOf[PersonCheck]), + withViolations = Seq( + WithViolation("id", "must not be empty", "", classOf[WithPersonCheck], toTest) + ) + ) + + // multiple groups with PersonCheck group enabled + assertViolations( + obj = toTest, + groups = Seq(classOf[OtherCheck], classOf[PersonCheck]), + withViolations = Seq( + WithViolation("id", "must not be empty", "", classOf[WithPersonCheck], toTest) + ) + ) + } + + test("ScalaValidator#isCascaded method validation - defined fields") { + // nested method validation fields + val owner = Person(id = "", name = "A. Einstein", address = DefaultAddress) + val car = Car( + id = 1234, + make = CarMake.Volkswagen, + model = "Beetle", + year = 1970, + owners = Seq(Person(id = "9999", name = "", address = DefaultAddress)), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyEnd = Some(LocalDate.now.plusYears(3)), + passengers = Seq( + Person(id = "1001", name = "R. Franklin", address = DefaultAddress) + ) + ) + val driver = Driver( + person = owner, + car = car + ) + + // no PersonCheck -- doesn't fail + assertViolations(obj = driver, groups = Seq(classOf[PersonCheck])) + // with default group - fail + assertViolations( + obj = driver, + withViolations = Seq( + WithViolation("car.owners[0].name", "must not be empty", "", classOf[Driver], driver), + WithViolation("car.validateId", "id may not be even", car, classOf[Driver], driver), + WithViolation( + "car.warrantyTimeValid.warrantyEnd", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Driver], + driver), + WithViolation( + "car.warrantyTimeValid.warrantyStart", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Driver], + driver), + WithViolation( + "car.year", + "must be greater than or equal to 2000", + 1970, + classOf[Driver], + driver), + WithViolation("person.id", "must not be empty", "", classOf[Driver], driver) + ) + ) + + // Validation with the default group does not return any violations with the standard + // Hibernate Validator because Scala only places annotations the executable fields of the + // case class by default and does copy them to the generated fields with the @field meta annotation. + Validation + .buildDefaultValidatorFactory() + .getValidator + .validate(driver) + .asScala + .isEmpty should be(true) + + val carWithInvalidNumberOfPassengers = Car( + id = 1235, + make = CarMake.Audi, + model = "A4", + year = 2010, + owners = Seq(Person(id = "9999", name = "M Bailey", address = DefaultAddress)), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyStart = Some(LocalDate.now), + warrantyEnd = Some(LocalDate.now.plusYears(3)), + passengers = Seq.empty + ) + val anotherDriver = Driver( + person = Person(id = "55555", name = "K. Brown", address = DefaultAddress), + car = carWithInvalidNumberOfPassengers + ) + + // ensure the class-level Car validation reports the correct property path of the field name + assertViolations( + obj = anotherDriver, + withViolations = Seq( + WithViolation( + "car", + "number of passenger(s) is not valid", + carWithInvalidNumberOfPassengers, + classOf[Driver], + anotherDriver) + ) + ) + } + + test("ScalaValidator#class-level annotations") { + // The Car case class is annotated with @ValidPassengerCount which checks if the number of + // passengers is greater than zero and less then some max specified in the annotation. + // In this test, no passengers are defined so the validation fails the class-level constraint. + // see: https://docs.jboss.org/hibernate/stable/scalaValidator/reference/en-US/html_single/#scalaValidator-usingscalaValidator-classlevel + val car = Car( + id = 1234, + make = CarMake.Volkswagen, + model = "Beetle", + year = 2004, + owners = Seq(Person(id = "9999", name = "K. Ann", address = DefaultAddress)), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyEnd = Some(LocalDate.now.plusYears(3)) + ) + + assertViolations( + obj = car, + withViolations = Seq( + WithViolation("", "number of passenger(s) is not valid", car, classOf[Car], car), + WithViolation("validateId", "id may not be even", car, classOf[Car], car), + WithViolation( + "warrantyTimeValid.warrantyEnd", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car), + WithViolation( + "warrantyTimeValid.warrantyStart", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car) + ) + ) + + // compare the class-level violation which should be the same as from the ScalaValidator + val fromUnderlyingViolations = validator.underlying.validate(car).asScala.toSet + fromUnderlyingViolations.size should equal(1) + val violation = fromUnderlyingViolations.head + violation.getPropertyPath.toString should equal("") + violation.getMessage should equal("number of passenger(s) is not valid") + violation.getRootBeanClass should equal(classOf[Car]) + violation.getInvalidValue should equal(car) + violation.getRootBean should equal(car) + + // Validation with the default group does not return any violations with the standard + // Hibernate Validator because Scala only places annotations the executable fields of the + // case class by default and does copy them to the generated fields with the @field meta annotation. + // However, class-level annotations will work. + val configuration = Validation + .byDefaultProvider() + .configure().asInstanceOf[HibernateValidatorConfiguration] + val mapping = configuration.createConstraintMapping() + mapping + .constraintDefinition(classOf[ValidPassengerCount]) + .validatedBy(classOf[ValidPassengerCountConstraintValidator]) + configuration.addMapping(mapping) + + val hibernateViolations = configuration.buildValidatorFactory.getValidator + .validate(car) + .asScala + // only has class-level constraint violation + val classLevelConstraintViolation = hibernateViolations.head + classLevelConstraintViolation.getPropertyPath.toString should be("") + classLevelConstraintViolation.getMessage should equal("number of passenger(s) is not valid") + classLevelConstraintViolation.getRootBeanClass should equal(classOf[Car]) + classLevelConstraintViolation.getRootBean == car should be(true) + + } + + test("ScalaValidator#cascaded validations") { + val validUser: User = User("1234567", "ion", "Other") + val invalidUser: User = User("", "anion", "M") + val nestedValidUser: Users = Users(Seq(validUser)) + val nestedInvalidUser: Users = Users(Seq(invalidUser)) + val nestedDuplicateUser: Users = Users(Seq(validUser, validUser)) + + assertViolations(obj = validUser) + assertViolations( + obj = invalidUser, + withViolations = Seq( + WithViolation("id", "must not be empty", "", classOf[User], invalidUser) + ) + ) + + assertViolations(obj = nestedValidUser) + assertViolations( + obj = nestedInvalidUser, + withViolations = Seq( + WithViolation("users[0].id", "must not be empty", "", classOf[Users], nestedInvalidUser)) + ) + + assertViolations(obj = nestedDuplicateUser) + } + + /* + * Builder withDescriptorCacheSize(...) tests + */ + test("ScalaValidator#withDescriptorCacheSize should override the default cache size") { + val customizedCacheSize: Long = 512 + val scalaValidator = ScalaValidator.builder + .withDescriptorCacheSize(customizedCacheSize) + scalaValidator.descriptorCacheSize should be(customizedCacheSize) + } + + test("ScalaValidator#validate is valid") { + val testUser: User = User(id = "9999", name = "April", gender = "F") + validator.validate(testUser).isEmpty should be(true) + } + + test("ScalaValidator#validate returns valid result even when case class has other fields") { + val testPerson: Person = Person(id = "9999", name = "April", address = DefaultAddress) + validator.validate(testPerson).isEmpty should be(true) + } + + test("ScalaValidator#validate is invalid") { + val testUser: User = User(id = "", name = "April", gender = "F") + val fieldAnnotation = getValidationAnnotations(classOf[User], "id").headOption + + assertViolations( + obj = testUser, + withViolations = Seq( + WithViolation("id", "must not be empty", "", classOf[User], testUser, fieldAnnotation) + ) + ) + } + + test( + "ScalaValidator#validate returns invalid result of both field validations and method validations") { + val testUser: User = User(id = "", name = "", gender = "Female") + + val idAnnotation: Option[Annotation] = + getValidationAnnotations(classOf[User], "id").headOption + val genderAnnotation: Option[Annotation] = + getValidationAnnotations(classOf[User], "gender").headOption + val methodAnnotation: Option[Annotation] = + getValidationAnnotations(classOf[User], "nameCheck").headOption + + assertViolations( + obj = testUser, + withViolations = Seq( + WithViolation( + "gender", + "Female not one of [F, M, Other]", + "Female", + classOf[User], + testUser, + genderAnnotation), + WithViolation("id", "must not be empty", "", classOf[User], testUser, idAnnotation), + WithViolation( + "nameCheck.name", + "cannot be empty", + testUser, + classOf[User], + testUser, + methodAnnotation) + ) + ) + } + + test("ScalaValidator#should register the user defined constraint validator") { + val testState: StateValidationExample = StateValidationExample(state = "NY") + val fieldAnnotation: Option[Annotation] = + getValidationAnnotations(classOf[StateValidationExample], "state").headOption + + assertViolations( + obj = testState, + withViolations = Seq( + WithViolation( + "state", + "Please register with state CA", + "NY", + classOf[StateValidationExample], + testState, + fieldAnnotation) + ) + ) + } + + test("ScalaValidator#secondary case class constructors") { + // the framework does not validate on construction, however, it must use + // the case class executable for finding the field validation annotations + // as Scala carries case class annotations on the executable for fields specified + // and annotated in the executable. Validations only apply to executable parameters that are + // also class fields, e.g., from the primary executable. + assertViolations(obj = TestJsonCreator("42")) + assertViolations(obj = TestJsonCreator(42)) + assertViolations(obj = TestJsonCreator2(scala.collection.immutable.Seq("1", "2", "3"))) + assertViolations( + obj = TestJsonCreator2(scala.collection.immutable.Seq(1, 2, 3), default = "Goodbye, world")) + + assertViolations(obj = TestJsonCreatorWithValidation("42")) + intercept[NumberFormatException] { + // can't validate after the fact -- the instance is already constructed, then we validate + assertViolations(obj = TestJsonCreatorWithValidation("")) + } + assertViolations(obj = TestJsonCreatorWithValidation(42)) + + // annotations are on primary executable + assertViolations( + obj = TestJsonCreatorWithValidations("99"), + withViolations = Seq( + WithViolation( + "int", + "99 not one of [42, 137]", + 99, + classOf[TestJsonCreatorWithValidations], + TestJsonCreatorWithValidations("99") + ) + ) + ) + assertViolations(obj = TestJsonCreatorWithValidations(42)) + + assertViolations(obj = new CaseClassWithMultipleConstructors("10001", "20002", "30003")) + assertViolations(obj = CaseClassWithMultipleConstructors(10001L, 20002L, 30003L)) + + assertViolations(obj = CaseClassWithMultipleConstructorsAnnotated(10001L, 20002L, 30003L)) + assertViolations(obj = + new CaseClassWithMultipleConstructorsAnnotated("10001", "20002", "30003")) + + assertViolations( + obj = CaseClassWithMultipleConstructorsAnnotatedAndValidations( + 10001L, + 20002L, + JUUID.randomUUID().toString)) + + assertViolations( + obj = new CaseClassWithMultipleConstructorsAnnotatedAndValidations( + "10001", + "20002", + JUUID.randomUUID().toString)) + + val invalid = new CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations( + "9999", + "10001", + JUUID.randomUUID().toString) + assertViolations( + obj = invalid, + withViolations = Seq( + WithViolation( + "number1", + "must be greater than or equal to 10000", + 9999, + classOf[CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations], + invalid + ) + ) + ) + + val d = InvalidDoublePerson("Andrea")("") + assertViolations(d) + val d2 = DoublePerson("Andrea")("") + assertViolations( + obj = d2, + withViolations = + Seq(WithViolation("otherName", "must not be empty", "", classOf[DoublePerson], d2))) + val d3 = ValidDoublePerson("Andrea")("") + val ve = intercept[ValidationException] { + validator.verify(d3) + } + ve.isInstanceOf[ConstraintViolationException] should be(true) + val cve = ve.asInstanceOf[ConstraintViolationException] + cve.getConstraintViolations.size should equal(2) + cve.getMessage should equal( + "\nValidation Errors:\t\tcheckOtherName: otherName must be longer than 3 chars, otherName: must not be empty\n\n") + + val d4 = PossiblyValidDoublePerson("Andrea")("") + assertViolations( + obj = d4, + withViolations = Seq( + WithViolation("otherName", "must not be empty", "", classOf[PossiblyValidDoublePerson], d4)) + ) + + val d5 = WithFinalValField( + JUUID.randomUUID().toString, + Some("Joan"), + Some("Jett"), + Some(Set(1, 2, 3)) + ) + assertViolations(d5) + } + + test("ScalaValidator#cycles") { + validator.validate(A("5678", B("9876", C("444", null)))).isEmpty should be(true) + validator.validate(D("1", E("2", F("3", None)))).isEmpty should be(true) + validator.validate(G("1", H("2", I("3", Seq.empty[G])))).isEmpty should be(true) + } + + test("ScalaValidator#scala BeanProperty") { + val value = AnnotatedBeanProperties(3) + value.setField1("") + assertViolations( + obj = value, + withViolations = Seq( + WithViolation("field1", "must not be empty", "", classOf[AnnotatedBeanProperties], value)) + ) + } + + test("ScalaValidator#case class with no executable params") { + val value = NoConstructorParams() + value.setId("1234") + assertViolations(value) + + assertViolations( + obj = NoConstructorParams(), + withViolations = Seq( + WithViolation( + "id", + "must not be empty", + "", + classOf[NoConstructorParams], + NoConstructorParams())) + ) + } + + test("ScalaValidator#case class annotated non executable field") { + assertViolations( + obj = AnnotatedInternalFields("1234", "thisisabigstring"), + withViolations = Seq( + WithViolation( + "company", + "must not be empty", + "", + classOf[AnnotatedInternalFields], + AnnotatedInternalFields("1234", "thisisabigstring")) + ) + ) + } + + test("ScalaValidator#inherited validation annotations") { + assertViolations( + obj = ImplementsAncestor(""), + withViolations = Seq( + WithViolation( + "field1", + "must not be empty", + "", + classOf[ImplementsAncestor], + ImplementsAncestor("")), + WithViolation( + "validateField1.field1", + "not a double value", + ImplementsAncestor(""), + classOf[ImplementsAncestor], + ImplementsAncestor("")) + ) + ) + + assertViolations( + obj = ImplementsAncestor("blimey"), + withViolations = Seq( + WithViolation( + "validateField1.field1", + "not a double value", + ImplementsAncestor("blimey"), + classOf[ImplementsAncestor], + ImplementsAncestor("blimey") + ) + ) + ) + + assertViolations(obj = ImplementsAncestor("3.141592653589793d")) + } + + test("ScalaValidator#getConstraintsForClass") { + // there is no resolution of nested types in a description + val driverCaseClassDescriptor = validator.getConstraintsForClass(classOf[Driver]) + driverCaseClassDescriptor.members.size should equal(2) + + verifyDriverClassDescriptor(driverCaseClassDescriptor) + } + + test("ScalaValidator#annotations") { + val caseClassDescriptor = validator.descriptorFactory.describe(classOf[TestAnnotations]) + caseClassDescriptor.members.size should equal(1) + val (name, memberDescriptor) = caseClassDescriptor.members.head + name should equal("cars") + memberDescriptor.isInstanceOf[PropertyDescriptor] should be(true) + memberDescriptor.asInstanceOf[PropertyDescriptor].annotations.length should equal(1) + memberDescriptor + .asInstanceOf[PropertyDescriptor].annotations.head.annotationType() should be( + classOf[NotEmpty]) + } + + test("ScalaValidator#generics 1") { + val caseClassDescriptor = + validator.descriptorFactory.describe(classOf[GenericTestCaseClass[_]]) + caseClassDescriptor.members.size should equal(1) + val value = GenericTestCaseClass(data = "Hello, World") + val results = validator.validate(value) + results.isEmpty should be(true) + + intercept[UnexpectedTypeException] { + validator.validate(GenericTestCaseClass(3)) + } + + val genericValue = GenericTestCaseClass[Seq[String]](Seq.empty[String]) + val violations = validator.validate(genericValue) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("data") + violation.getRootBeanClass should equal(classOf[GenericTestCaseClass[Seq[String]]]) + violation.getLeafBean should equal(genericValue) + violation.getInvalidValue should equal(Seq.empty[String]) + + validator.validate(GenericMinTestCaseClass(5)).isEmpty should be(true) + + val caseClassDescriptor1 = + validator.descriptorFactory.describe(classOf[GenericTestCaseClassMultipleTypes[_, _, _]]) + caseClassDescriptor1.members.size should equal(3) + + val value1 = GenericTestCaseClassMultipleTypes( + data = None, + things = Seq(1, 2, 3), + otherThings = Seq( + UUIDExample("1234"), + UUIDExample(JUUID.randomUUID().toString), + UUIDExample(JUUID.randomUUID().toString)) + ) + val results1 = validator.validate(value1) + results1.isEmpty should be(true) + } + + test("ScalaValidator#generics 2") { + val value: Page[Person] = Page[Person]( + List( + Person(id = "9999", name = "April", address = DefaultAddress), + Person(id = "9999", name = "April", address = DefaultAddress) + ), + 0, + None, + None) + + val caseClassDescriptor = + validator.descriptorFactory.describe(classOf[Page[Person]]) + caseClassDescriptor.members.size should equal(4) + + val violations = validator.validate(value) + violations.size should equal(1) + + val violation = violations.head + violation.getPropertyPath.toString should be("pageSize") + violation.getMessage should be("must be greater than or equal to 1") + violation.getRootBeanClass should equal(classOf[Page[Person]]) + violation.getRootBean should equal(value) + violation.getInvalidValue should equal(0) + } + + test("ScalaValidator#test boxed primitives") { + validator + .validateValue( + classOf[CaseClassWithBoxedPrimitives], + "events", + java.lang.Integer.valueOf(42)).isEmpty should be(true) + } + + test("ScalaValidator#find executable 1") { + /* CaseClassWithMultipleConstructorsAnnotatedAndValidations */ + val clazzDescriptor: ClassDescriptor = + Reflector + .describe(classOf[CaseClassWithMultipleConstructorsAnnotatedAndValidations]) + .asInstanceOf[ClassDescriptor] + + val constructorDescriptor: ConstructorDescriptor = + validator.descriptorFactory.findConstructor(clazzDescriptor /*, Seq.empty*/ ) + + // by default this will find Seq(Long, Long, String) + val constructorParameterTypes: Array[Class[_]] = + constructorDescriptor.constructor.getParameterTypes() + constructorParameterTypes.length should equal(3) + constructorParameterTypes(0) should equal(classOf[Long]) + constructorParameterTypes(1) should equal(classOf[Long]) + constructorParameterTypes(2) should equal(classOf[String]) + + val caseClazzDescriptor = validator.getConstraintsForClass( + classOf[CaseClassWithMultipleConstructorsAnnotatedAndValidations]) + caseClazzDescriptor.members.size should equal(3) + caseClazzDescriptor.members.contains("number1") should be(true) + caseClazzDescriptor.members("number1").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("number2") should be(true) + caseClazzDescriptor.members("number2").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("uuid") should be(true) + caseClazzDescriptor.members("uuid").annotations.isEmpty should be(true) + } + + test("ScalaValidator#find executable 2") { + /* CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations */ + val clazzDescriptor = Reflector + .describe(classOf[CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations]) + .asInstanceOf[ClassDescriptor] + + val constructorDescriptor = + validator.descriptorFactory.findConstructor(clazzDescriptor) + + // by default this will find Seq(Long, Long, String) + val constructorParameterTypes: Array[Class[_]] = + constructorDescriptor.constructor.getParameterTypes() + constructorParameterTypes.length should equal(3) + constructorParameterTypes(0) should equal(classOf[Long]) + constructorParameterTypes(1) should equal(classOf[Long]) + constructorParameterTypes(2) should equal(classOf[String]) + + val caseClazzDescriptor = validator.getConstraintsForClass( + classOf[CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations]) + caseClazzDescriptor.members.size should equal(3) + caseClazzDescriptor.members.contains("number1") should be(true) + caseClazzDescriptor.members("number1").annotations.length should equal(1) + caseClazzDescriptor.members.contains("number2") should be(true) + caseClazzDescriptor.members("number2").annotations.length should equal(1) + caseClazzDescriptor.members.contains("uuid") should be(true) + caseClazzDescriptor.members("uuid").annotations.length should equal(1) + } + + test("ScalaValidator#find executable 3") { + /* CaseClassWithMultipleConstructorsAnnotated */ + val clazzDescriptor = Reflector + .describe(classOf[CaseClassWithMultipleConstructorsAnnotated]) + .asInstanceOf[ClassDescriptor] + + val constructorDescriptor = + validator.descriptorFactory.findConstructor(clazzDescriptor) + + // this should find Seq(Long, Long, Long) + val constructorParameterTypes: Array[Class[_]] = + constructorDescriptor.constructor.getParameterTypes() + constructorParameterTypes.length should equal(3) + constructorParameterTypes(0) should equal(classOf[Long]) + constructorParameterTypes(1) should equal(classOf[Long]) + constructorParameterTypes(2) should equal(classOf[Long]) + + val caseClazzDescriptor = + validator.getConstraintsForClass(classOf[CaseClassWithMultipleConstructorsAnnotated]) + caseClazzDescriptor.members.size should equal(3) + caseClazzDescriptor.members.contains("number1") should be(true) + caseClazzDescriptor.members("number1").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("number2") should be(true) + caseClazzDescriptor.members("number2").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("number3") should be(true) + caseClazzDescriptor.members("number3").annotations.isEmpty should be(true) + } + + test("ScalaValidator#find executable 4") { + /* CaseClassWithMultipleConstructors */ + val clazzDescriptor = Reflector + .describe(classOf[CaseClassWithMultipleConstructors]) + .asInstanceOf[ClassDescriptor] + + val constructorDescriptor = + validator.descriptorFactory.findConstructor(clazzDescriptor) + + // this should find Seq(Long, Long, Long) + val constructorParameterTypes: Array[Class[_]] = + constructorDescriptor.constructor.getParameterTypes() + constructorParameterTypes.length should equal(3) + constructorParameterTypes(0) should equal(classOf[Long]) + constructorParameterTypes(1) should equal(classOf[Long]) + constructorParameterTypes(2) should equal(classOf[Long]) + + val caseClazzDescriptor = + validator.getConstraintsForClass(classOf[CaseClassWithMultipleConstructors]) + caseClazzDescriptor.members.size should equal(3) + caseClazzDescriptor.members.contains("number1") should be(true) + caseClazzDescriptor.members("number1").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("number2") should be(true) + caseClazzDescriptor.members("number2").annotations.isEmpty should be(true) + caseClazzDescriptor.members.contains("number3") should be(true) + caseClazzDescriptor.members("number3").annotations.isEmpty should be(true) + } + + test("ScalaValidator#find executable 5") { + /* TestJsonCreatorWithValidations */ + val clazzDescriptor = Reflector + .describe(classOf[TestJsonCreatorWithValidations]) + .asInstanceOf[ClassDescriptor] + + val constructorDescriptor = + validator.descriptorFactory.findConstructor(clazzDescriptor) + + // should find Seq(int) + val constructorParameterTypes: Array[Class[_]] = + constructorDescriptor.constructor.getParameterTypes() + constructorParameterTypes.length should equal(1) + constructorParameterTypes.head should equal(classOf[Int]) + + val caseClazzDescriptor = + validator.getConstraintsForClass(classOf[TestJsonCreatorWithValidations]) + caseClazzDescriptor.members.size should equal(1) + caseClazzDescriptor.members.contains("int") should be(true) + caseClazzDescriptor.members("int").annotations.length should equal(1) + } + + test("ScalaValidator#with Logging trait") { + val value = PersonWithLogging(42, "Baz Var", None) + validator.validate(value).isEmpty should be(true) + } + + private[this] def newJavaSet(numElements: Int): util.Set[Int] = { + val result: util.Set[Int] = new util.HashSet[Int]() + for (i <- 1 to numElements) { + result.add(i) + } + result + } + + private[this] def verifyDriverClassDescriptor( + driverCaseClassDescriptor: CaseClassDescriptor[Driver] + ): Unit = { + val person = driverCaseClassDescriptor.members("person") + person.scalaType.erasure should equal(classOf[Person]) + // cascaded types will have a contained scala type + person.cascadedScalaType.isDefined should be(true) + person.cascadedScalaType.get.erasure should equal(classOf[Person]) + person.annotations.isEmpty should be(true) + person.isCascaded should be(true) + + // person is cascaded, so let's look it up + val personDescriptor = validator.getConstraintsForClass(classOf[Person]) + personDescriptor.members.size should equal(6) + val city = personDescriptor.members("city") + city.scalaType.erasure should equal(classOf[String]) + city.annotations.isEmpty should be(true) + city.isCascaded should be(false) + + val name = personDescriptor.members("name") + name.scalaType.erasure == classOf[String] should be(true) + name.annotations.length should equal(1) + name.annotations.head.annotationType() should equal(classOf[NotEmpty]) + name.isCascaded should be(false) + + val state = personDescriptor.members("state") + state.scalaType.erasure should equal(classOf[String]) + state.annotations.isEmpty should be(true) + state.isCascaded should be(false) + + val company = personDescriptor.members("company") + company.scalaType.erasure should equal(classOf[String]) + company.annotations.isEmpty should be(true) + company.isCascaded should be(false) + + val id = personDescriptor.members("id") + id.scalaType.erasure == classOf[String] should be(true) + id.annotations.length should equal(1) + id.annotations.head.annotationType() should equal(classOf[NotEmpty]) + id.isCascaded should be(false) + + val address = personDescriptor.members("address") + address.scalaType.erasure should equal(classOf[Address]) + address.annotations.isEmpty should be(true) + address.isCascaded should be(true) + + // address is cascaded, so let's look it up + val addressDescriptor = validator.getConstraintsForClass(classOf[Address]) + addressDescriptor.members.size should equal(6) + + val line1 = addressDescriptor.members("line1") + line1.scalaType.erasure == classOf[String] should be(true) + line1.annotations.isEmpty should be(true) + line1.isCascaded should be(false) + + val line2 = addressDescriptor.members("line2") + line2.scalaType.erasure == classOf[Option[String]] should be(true) + line2.annotations.isEmpty should be(true) + line2.isCascaded should be(false) + + val addressCity = addressDescriptor.members("city") + addressCity.scalaType.erasure == classOf[String] should be(true) + addressCity.annotations.isEmpty should be(true) + addressCity.isCascaded should be(false) + + val addressState = addressDescriptor.members("state") + addressState.scalaType.erasure == classOf[String] should be(true) + addressState.annotations.length should equal(1) + addressState.annotations.head.annotationType() should equal(classOf[StateConstraint]) + addressState.isCascaded should be(false) + + val zipcode = addressDescriptor.members("zipcode") + zipcode.scalaType.erasure == classOf[String] should be(true) + zipcode.annotations.isEmpty should be(true) + zipcode.isCascaded should be(false) + + val additionalPostalCode = addressDescriptor.members("additionalPostalCode") + additionalPostalCode.scalaType.erasure == classOf[Option[String]] should be(true) + // not cascaded so no contained type + additionalPostalCode.cascadedScalaType.isEmpty should be(true) + additionalPostalCode.annotations.isEmpty should be(true) + additionalPostalCode.isCascaded should be(false) + + personDescriptor.methods.length should equal(1) + personDescriptor.methods.head.method.getName should equal("validateAddress") + personDescriptor.methods.head.annotations.length should equal(1) + personDescriptor.methods.head.annotations.head.annotationType() should equal( + classOf[MethodValidation]) + personDescriptor.methods.head.members.size should equal(0) + + val car = driverCaseClassDescriptor.members("car") + car.scalaType.erasure should equal(classOf[Car]) + car.cascadedScalaType.isDefined should be(true) + car.cascadedScalaType.get.erasure should equal(classOf[Car]) + car.annotations.isEmpty should be(true) + car.isCascaded should be(true) + + // car is cascaded, so let's look it up + val carDescriptor = validator.getConstraintsForClass(classOf[Car]) + carDescriptor.members.size should equal(13) + + // lets just look at annotated and cascaded fields in car + val year = carDescriptor.members("year") + year.scalaType.erasure == classOf[Int] should be(true) + year.annotations.length should equal(1) + year.annotations.head.annotationType() should equal(classOf[Min]) + year.isCascaded should be(false) + + val owners = carDescriptor.members("owners") + owners.scalaType.erasure == classOf[Seq[Person]] should be(true) + owners.cascadedScalaType.isDefined should be(true) + owners.cascadedScalaType.get.erasure should equal(classOf[Person]) + owners.annotations.isEmpty should be(true) + owners.isCascaded should be(true) + + val licensePlate = carDescriptor.members("licensePlate") + licensePlate.scalaType.erasure == classOf[String] should be(true) + licensePlate.annotations.length should equal(2) + licensePlate.annotations.head.annotationType() should equal(classOf[NotEmpty]) + licensePlate.annotations.last.annotationType() should equal(classOf[Size]) + licensePlate.isCascaded should be(false) + + val numDoors = carDescriptor.members("numDoors") + numDoors.scalaType.erasure == classOf[Int] should be(true) + numDoors.annotations.length should equal(1) + numDoors.annotations.head.annotationType() should equal(classOf[Min]) + numDoors.isCascaded should be(false) + + val passengers = carDescriptor.members("passengers") + passengers.scalaType.erasure == classOf[Seq[Person]] should be(true) + passengers.cascadedScalaType.isDefined should be(true) + passengers.cascadedScalaType.get.erasure should equal(classOf[Person]) + passengers.annotations.isEmpty should be(true) + passengers.isCascaded should be(true) + + carDescriptor.methods.length should equal(4) + val mappedCarMethods = carDescriptor.methods.map { methodDescriptor => + methodDescriptor.method.getName -> methodDescriptor + }.toMap + + val validateId = mappedCarMethods("validateId") + validateId.members.isEmpty should be(true) + validateId.annotations.length should equal(1) + validateId.annotations.head.annotationType() should equal(classOf[MethodValidation]) + validateId.members.size should equal(0) + + val validateYearBeforeNow = mappedCarMethods("validateYearBeforeNow") + validateYearBeforeNow.members.isEmpty should be(true) + validateYearBeforeNow.annotations.length should equal(1) + validateYearBeforeNow.annotations.head.annotationType() should equal(classOf[MethodValidation]) + validateYearBeforeNow.members.size should equal(0) + + val ownershipTimesValid = mappedCarMethods("ownershipTimesValid") + ownershipTimesValid.members.isEmpty should be(true) + ownershipTimesValid.annotations.length should equal(1) + ownershipTimesValid.annotations.head.annotationType() should equal(classOf[MethodValidation]) + ownershipTimesValid.members.size should equal(0) + + val warrantyTimeValid = mappedCarMethods("warrantyTimeValid") + warrantyTimeValid.members.isEmpty should be(true) + warrantyTimeValid.annotations.length should equal(1) + warrantyTimeValid.annotations.head.annotationType() should equal(classOf[MethodValidation]) + warrantyTimeValid.members.size should equal(0) + } + + test("ScalaValidator#class with incorrect method validation definition 1") { + val value = WithIncorrectlyDefinedMethodValidation("abcd1234") + + val e = intercept[ConstraintDeclarationException] { + validator.validate(value) + } + e.getMessage should equal( + "Methods annotated with @MethodValidation must not declare any arguments") + } + + test("ScalaValidator#class with incorrect method validation definition 2") { + val value = AnotherIncorrectlyDefinedMethodValidation("abcd1234") + + val e = intercept[ConstraintDeclarationException] { + validator.validate(value) + } + e.getMessage should equal( + "Methods annotated with @MethodValidation must return a com.twitter.util.validation.engine.MethodValidationResult") + } + + private[this] def getValidationAnnotations( + clazz: Class[_], + name: String + ): Array[Annotation] = { + val caseClassDescriptor = validator.descriptorFactory.describe(clazz) + caseClassDescriptor.members.get(name) match { + case Some(property: PropertyDescriptor) => + property.annotations + case _ => Array.empty[Annotation] + } + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/WrongMessageInterpolator.scala b/util-validator/src/test/scala/com/twitter/util/validation/WrongMessageInterpolator.scala new file mode 100644 index 0000000000..6691819240 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/WrongMessageInterpolator.scala @@ -0,0 +1,17 @@ +package com.twitter.util.validation + +import jakarta.validation.MessageInterpolator +import java.util.Locale + +class WrongMessageInterpolator extends MessageInterpolator { + override def interpolate( + s: String, + context: MessageInterpolator.Context + ): String = "Whatever you entered was wrong" + + override def interpolate( + s: String, + context: MessageInterpolator.Context, + locale: Locale + ): String = "Whatever you entered was wrong" +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/caseclasses.scala b/util-validator/src/test/scala/com/twitter/util/validation/caseclasses.scala new file mode 100644 index 0000000000..d35de93d9d --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/caseclasses.scala @@ -0,0 +1,537 @@ +package com.twitter.util.validation + +import com.fasterxml.jackson.annotation.JsonCreator +import com.twitter.util.Try +import com.twitter.util.logging.Logging +import com.twitter.util.validation.constraints._ +import com.twitter.util.validation.engine.MethodValidationResult +import jakarta.validation.{Payload, Valid, ValidationException} +import jakarta.validation.constraints._ +import java.time.LocalDate +import scala.annotation.meta.field +import scala.beans.BeanProperty + +object OuterObject { + + case class NestedCaseClass(id: Long) + + object InnerObject { + + case class SuperNestedCaseClass(@Min(5) id: Long) + } +} + +object caseclasses { + + object Path { + val FieldSeparator = "." + + /** An empty path (no names) */ + val Empty: Path = Path(Seq.empty[String]) + + /** Creates a new path from the given name */ + def apply(name: String): Path = new Path(Seq(name)) + } + + /** + * Represents the navigation path from an object to another in an object graph. + */ + case class Path(names: Seq[String]) + + case class Page[T]( + data: List[T], + @Min(1) pageSize: Int, + next: Option[Long], + previous: Option[Long]) + + case class UnicodeNameCaseClass(@Max(3) `winning-id`: Int, @NotEmpty name: String) + + class EvaluateConstraints + + trait NotACaseClassTrait { + @UUID def uuid: String + } + + class NotACaseClass( + @NotNull manufacturer: String, + @NotNull @Size(min = 2, max = 14) licensePlate: String, + @Min(2) seatCount: Int, + uuid: String) + + final case class WithFinalValField( + @UUID uuid: String, + @NotEmpty first: Option[String], + last: Option[String], + private val permissions: Option[Set[Int]]) + + case class WithBogusCascadedField( + @NotNull @Size(min = 2, max = 14) licensePlate: String, + @Valid ints: Seq[Int]) + + case class WithMetaFieldAnnotation(@(NotEmpty @field) id: String) + + // @NotEmpty is copied to generated field, but @Size is not. We should find both. + case class WithPartialMetaFieldAnnotation(@(NotEmpty @field) @Size(min = 2, max = 10) id: String) + + case class PathNotEmpty(@NotEmpty path: Path, @NotEmpty id: String) + + case class Customer(@NotEmpty first: String, last: String) + + case class AlwaysFails(@InvalidConstraint name: String) + + case class AlwaysFailsMethodValidation(@NotEmpty id: String) { + @MethodValidation(fields = Array("id")) + def checkId: MethodValidationResult = { + throw new ValidationException("oh noes!") + } + } + + // throws a ConstraintDefinitionException as `@ethodValidation` must not define args + case class WithIncorrectlyDefinedMethodValidation(@NotEmpty id: String) { + @MethodValidation(fields = Array("id")) + def checkId(a: String): MethodValidationResult = { + MethodValidationResult.validIfTrue(a == id, "id not equal") + } + } + + // throws a ConstraintDefinitionException as `@ethodValidation` must return a MethodValidationResult + case class AnotherIncorrectlyDefinedMethodValidation(@NotEmpty id: String) { + @MethodValidation(fields = Array("id")) + def checkId: Boolean = id.nonEmpty + } + + case class SmallCar( + @NotNull manufacturer: String, + @NotNull @Size(min = 2, max = 14) licensePlate: String, + @Min(2) seatCount: Int) + + trait RentalStationMixin { + @Min(10) def name: String + } + + // doesn't define any field + trait RandoMixin + + case class RentalStation(@NotNull name: String) { + def rentCar( + @NotNull customer: Customer, + @NotNull @Future start: LocalDate, + @Min(1) duration: Int + ): Unit = ??? + + def updateCustomerRecords( + @NotEmpty @Size(min = 1) @Valid customers: Seq[Customer] + ): Unit = ??? + + @NotEmpty + @Size(min = 1) + def getCustomers: Seq[Customer] = ??? + + def listCars(@NotEmpty `airport-code`: String): Seq[String] = ??? + + @ConsistentDateParameters + def reserve( + start: LocalDate, + end: LocalDate + ): Boolean = ??? + } + + trait OtherCheck + trait PersonCheck + + case class MinIntExample(@Min(1) numberValue: Int) + + // CountryCode + case class CountryCodeExample(@CountryCode countryCode: String) + case class CountryCodeOptionExample(@CountryCode countryCode: Option[String]) + case class CountryCodeSeqExample(@CountryCode countryCode: Seq[String]) + case class CountryCodeOptionSeqExample(@CountryCode countryCode: Option[Seq[String]]) + case class CountryCodeArrayExample(@CountryCode countryCode: Array[String]) + case class CountryCodeOptionArrayExample(@CountryCode countryCode: Option[Array[String]]) + case class CountryCodeInvalidTypeExample(@CountryCode countryCode: Long) + case class CountryCodeOptionInvalidTypeExample(@CountryCode countryCode: Option[Long]) + // OneOf + case class OneOfExample(@OneOf(value = Array("a", "B", "c")) enumValue: String) + case class OneOfOptionExample(@OneOf(value = Array("a", "B", "c")) enumValue: Option[String]) + case class OneOfSeqExample(@OneOf(Array("a", "B", "c")) enumValue: Seq[String]) + case class OneOfOptionSeqExample(@OneOf(Array("a", "B", "c")) enumValue: Option[Seq[String]]) + case class OneOfInvalidTypeExample(@OneOf(Array("a", "B", "c")) enumValue: Long) + case class OneOfOptionInvalidTypeExample(@OneOf(Array("a", "B", "c")) enumValue: Option[Long]) + // UUID + case class UUIDExample(@UUID uuid: String) + case class UUIDOptionExample(@UUID uuid: Option[String]) + + case class NestedOptionExample( + @NotEmpty id: String, + name: String, + @Valid user: Option[User]) + + // multiple annotations + case class User( + @NotEmpty id: String, + name: String, + @OneOf(Array("F", "M", "Other")) gender: String) { + @MethodValidation(fields = Array("name")) + def nameCheck: MethodValidationResult = { + MethodValidationResult.validIfTrue(name.length > 0, "cannot be empty") + } + } + // nested NotEmpty on a Scala collection of a case class type + case class Users(@NotEmpty @Valid users: Seq[User]) + + // Other fields not in executable + case class Person(@NotEmpty id: String, @NotEmpty name: String, @Valid address: Address) { + val company: String = "Twitter" + val city: String = "San Francisco" + val state: String = "California" + + @MethodValidation(fields = Array("address")) + def validateAddress: MethodValidationResult = + MethodValidationResult.validIfTrue(address.state.nonEmpty, "state must be specified") + } + + case class WithPersonCheck( + @NotEmpty(groups = Array(classOf[PersonCheck])) id: String, + @NotEmpty name: String) + + case class NoConstructorParams() { + @BeanProperty @NotEmpty var id: String = "" + } + + case class AnnotatedInternalFields( + @NotEmpty id: String, + @Size(min = 1, max = 100) bigString: String) { + @NotEmpty val company: String = "" // this will be validated + } + + case class AnnotatedBeanProperties(@Min(1) numbers: Int) { + @BeanProperty @NotEmpty var field1: String = "default" + } + + case class Address( + line1: String, + line2: Option[String] = None, + city: String, + @StateConstraint(payload = Array(classOf[StateConstraintPayload])) state: String, + zipcode: String, + additionalPostalCode: Option[String] = None) + + // Customer defined validations + case class StateValidationExample(@StateConstraint state: String) + // Integration test + case class ValidateUserRequest( + @NotEmpty @Pattern(regexp = "[a-z]+") userName: String, + @Max(value = 9999) id: Long, + title: String) + + case class NestedUserPayload(id: String, job: String) extends Payload + case class NestedUser( + @NotEmpty id: String, + @Valid person: Person, + @OneOf(Array("F", "M", "Other")) gender: String, + job: String) { + @MethodValidation(fields = Array("job"), payload = Array(classOf[NestedUserPayload])) + def jobCheck: MethodValidationResult = { + MethodValidationResult.validIfTrue( + job.length > 0, + "cannot be empty", + payload = Some(NestedUserPayload(id, job)) + ) + } + } + + case class Division( + name: String, + @Valid team: Seq[Group]) + + case class Persons( + @NotEmpty id: String, + @Min(1) @Valid people: Seq[Person]) + + case class Group( + @NotEmpty id: String, + @Min(1) @Valid people: Seq[Person]) + + case class CustomerAccount( + @NotEmpty accountName: String, + @Valid customer: Option[User]) + + case class CaseClassWithBoxedPrimitives(events: Integer, errors: Integer) + + case class CollectionOfCollection( + @NotEmpty id: String, + @Valid @Min(1) people: Seq[Persons]) + + case class CollectionWithArray( + @NotEmpty @Size(min = 1, max = 30) names: Array[String], + @NotEmpty @Max(2) users: Array[User]) + + // cycle -- impossible to instantiate without using null to terminate somewhere + case class A(@NotEmpty id: String, @Valid b: B) + case class B(@NotEmpty id: String, @Valid c: C) + case class C(@NotEmpty id: String, @Valid a: A) + // another cycle -- fails as we detect the D type in F + case class D(@NotEmpty id: String, @Valid e: E) + case class E(@NotEmpty id: String, @Valid f: F) + case class F(@NotEmpty id: String, @Valid d: Option[D]) + // last one -- fails as we detect the G type in I + case class G(@NotEmpty id: String, @Valid h: H) + case class H(@NotEmpty id: String, @Valid i: I) + case class I(@NotEmpty id: String, @Valid g: Seq[G]) + + // no validation -- no annotations + object TestJsonCreator { + @JsonCreator + def apply(s: String): TestJsonCreator = TestJsonCreator(s.toInt) + } + case class TestJsonCreator(int: Int) + + // no validation -- no annotations + object TestJsonCreator2 { + @JsonCreator + def apply(strings: Seq[String]): TestJsonCreator2 = TestJsonCreator2(strings.map(_.toInt)) + } + case class TestJsonCreator2(ints: Seq[Int], default: String = "Hello, World") + + // no validation -- annotation is on a factory method "executable" + object TestJsonCreatorWithValidation { + @JsonCreator + def apply(@NotEmpty s: String): TestJsonCreatorWithValidation = + TestJsonCreatorWithValidation(s.toInt) + } + case class TestJsonCreatorWithValidation(int: Int) + + // this should validate per the primary executable annotation + object TestJsonCreatorWithValidations { + @JsonCreator + def apply(@NotEmpty s: String): TestJsonCreatorWithValidations = + TestJsonCreatorWithValidations(s.toInt) + } + case class TestJsonCreatorWithValidations(@OneOf(Array("42", "137")) int: Int) + + // no validation -- no annotations + case class CaseClassWithMultipleConstructors(number1: Long, number2: Long, number3: Long) { + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } + } + + // no validation -- no annotations + case class CaseClassWithMultipleConstructorsAnnotated( + number1: Long, + number2: Long, + number3: Long) { + @JsonCreator + def this(numberAsString1: String, numberAsString2: String, numberAsString3: String) { + this(numberAsString1.toLong, numberAsString2.toLong, numberAsString3.toLong) + } + } + + // no validation -- annotations are on secondary executable + case class CaseClassWithMultipleConstructorsAnnotatedAndValidations( + number1: Long, + number2: Long, + uuid: String) { + @JsonCreator + def this( + @NotEmpty numberAsString1: String, + @OneOf(Array("10001", "20002", "30003")) numberAsString2: String, + @UUID thirdArgument: String + ) { + this(numberAsString1.toLong, numberAsString2.toLong, thirdArgument) + } + } + + // this should validate -- annotations are on the primary executable + case class CaseClassWithMultipleConstructorsPrimaryAnnotatedAndValidations( + @Min(10000) number1: Long, + @OneOf(Array("10001", "20002", "30003")) number2: Long, + @UUID uuid: String) { + def this( + numberAsString1: String, + numberAsString2: String, + thirdArgument: String + ) { + this(numberAsString1.toLong, numberAsString2.toLong, thirdArgument) + } + } + + trait AncestorWithValidation { + @NotEmpty def field1: String + + @MethodValidation(fields = Array("field1")) + def validateField1: MethodValidationResult = + MethodValidationResult.validIfTrue(Try(field1.toDouble).isReturn, "not a double value") + } + + case class ImplementsAncestor(override val field1: String) extends AncestorWithValidation + + case class InvalidDoublePerson(@NotEmpty name: String)(@NotEmpty otherName: String) + case class DoublePerson(@NotEmpty name: String)(@NotEmpty val otherName: String) + case class ValidDoublePerson(@NotEmpty name: String)(@NotEmpty otherName: String) { + @MethodValidation + def checkOtherName: MethodValidationResult = + MethodValidationResult.validIfTrue( + otherName.length >= 3, + "otherName must be longer than 3 chars") + } + case class PossiblyValidDoublePerson(@NotEmpty name: String)(@NotEmpty otherName: String) { + def referenceSecondArgInMethod: String = + s"The otherName value is $otherName" + } + + case class CaseClassWithIntAndLocalDate( + @NotEmpty name: String, + age: Int, + age2: Int, + age3: Int, + dateTime: LocalDate, + dateTime2: LocalDate, + dateTime3: LocalDate, + dateTime4: LocalDate, + @NotEmpty dateTime5: Option[LocalDate]) + + case class SimpleCar( + id: Long, + make: CarMake, + model: String, + @Min(2000) year: Int, + @NotEmpty @Size(min = 2, max = 14) licensePlate: String, //multiple annotations + @Min(0) numDoors: Int = 4, + manual: Boolean = false) + + // executable annotation -- note the validator must specify a `@SupportedValidationTarget` of annotated element + case class CarWithPassengerCount @ValidPassengerCountReturnValue(max = 1) ( + passengers: Seq[Person]) + + // class-level annotation + @ValidPassengerCount(max = 4) + case class Car( + id: Long, + make: CarMake, + model: String, + @Min(2000) year: Int, + @Valid owners: Seq[Person], + @NotEmpty @Size(min = 2, max = 14) licensePlate: String, //multiple annotations + @Min(0) numDoors: Int = 4, + manual: Boolean = false, + ownershipStart: LocalDate = LocalDate.now, + ownershipEnd: LocalDate = LocalDate.now.plusYears(1), + warrantyStart: Option[LocalDate] = None, + warrantyEnd: Option[LocalDate] = None, + @Valid passengers: Seq[Person] = Seq.empty[Person]) { + + @MethodValidation + def validateId: MethodValidationResult = { + MethodValidationResult.validIfTrue(id % 2 == 1, "id may not be even") + } + + @MethodValidation + def validateYearBeforeNow: MethodValidationResult = { + val thisYear = LocalDate.now().getYear + val yearMoreThanOneYearInFuture: Boolean = + if (year > thisYear) { (year - thisYear) > 1 } + else false + MethodValidationResult.validIfFalse( + yearMoreThanOneYearInFuture, + "Model year can be at most one year newer." + ) + } + + @MethodValidation(fields = Array("ownershipEnd")) + def ownershipTimesValid: MethodValidationResult = { + validateTimeRange( + ownershipStart, + ownershipEnd, + "ownershipStart", + "ownershipEnd" + ) + } + + @MethodValidation(fields = Array("warrantyStart", "warrantyEnd")) + def warrantyTimeValid: MethodValidationResult = { + validateTimeRange( + warrantyStart, + warrantyEnd, + "warrantyStart", + "warrantyEnd" + ) + } + } + + case class Driver( + @Valid person: Person, + @Valid car: Car) + + case class MultipleConstraints( + @NotEmpty @Max(100) ints: Seq[Int], + @NotEmpty @Pattern(regexp = "\\d") digits: String, + @NotEmpty @Size(min = 3, max = 3) @OneOf(Array("how", "now", "brown", + "cow")) strings: Seq[String]) + + // maps not supported + case class MapsAndMaps( + @NotEmpty id: String, + @AssertTrue bool: Boolean, + @Valid carAndDriver: Map[Person, SimpleCar]) + + // @Valid only visible on field not on type arg + case class TestAnnotations( + @NotEmpty cars: Seq[SimpleCar @Valid]) + + case class GenericTestCaseClass[T](@NotEmpty @Valid data: T) + case class GenericMinTestCaseClass[T](@Min(4) data: T) + case class GenericTestCaseClassMultipleTypes[T, U, V]( + @NotEmpty @Valid data: T, + @Size(max = 100) things: Seq[U], + @Size(min = 3, max = 3) @Valid otherThings: Seq[V]) + case class GenericTestCaseClassWithMultipleArgs[T]( + @NotEmpty data: T, + @Min(5) number: Int) + + def validateTimeRange( + startTime: Option[LocalDate], + endTime: Option[LocalDate], + startTimeProperty: String, + endTimeProperty: String + ): MethodValidationResult = { + + val rangeDefined = startTime.isDefined && endTime.isDefined + val partialRange = !rangeDefined && (startTime.isDefined || endTime.isDefined) + + if (rangeDefined) + validateTimeRange(startTime.get, endTime.get, startTimeProperty, endTimeProperty) + else if (partialRange) + MethodValidationResult.Invalid( + "both %s and %s are required for a valid range".format(startTimeProperty, endTimeProperty) + ) + else MethodValidationResult.Valid + } + + def validateTimeRange( + startTime: LocalDate, + endTime: LocalDate, + startTimeProperty: String, + endTimeProperty: String + ): MethodValidationResult = { + + MethodValidationResult.validIfTrue( + startTime.isBefore(endTime), + "%s [%s] must be after %s [%s]" + .format(endTimeProperty, endTime.toEpochDay, startTimeProperty, startTime.toEpochDay) + ) + } + + case class PersonWithLogging( + id: Int, + name: String, + age: Option[Int], + age_with_default: Option[Int] = None, + nickname: String = "unknown") + extends Logging + + case class UsersRequest( + @Max(100) @DoesNothing max: Int, + @Past @DoesNothing startDate: Option[LocalDate], + @DoesNothing verbose: Boolean = false) +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/constraints/ConsistentDateParameters.java b/util-validator/src/test/scala/com/twitter/util/validation/constraints/ConsistentDateParameters.java new file mode 100644 index 0000000000..98a14c3865 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/constraints/ConsistentDateParameters.java @@ -0,0 +1,30 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.Documented; +import java.lang.annotation.Retention; +import java.lang.annotation.Target; + +import com.twitter.util.validation.internal.validators.ConsistentDateParametersValidator; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +import static java.lang.annotation.ElementType.ANNOTATION_TYPE; +import static java.lang.annotation.ElementType.METHOD; +import static java.lang.annotation.RetentionPolicy.RUNTIME; + +@Target({ METHOD, ANNOTATION_TYPE }) +@Retention(RUNTIME) +@Constraint(validatedBy = { ConsistentDateParametersValidator.class }) +@Documented +public @interface ConsistentDateParameters { + + /** message */ + String message() default "start is not before end"; + + /** groups */ + Class[] groups() default { }; + + /** payload */ + Class[] payload() default { }; +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraint.java b/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraint.java new file mode 100644 index 0000000000..946e4aa893 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraint.java @@ -0,0 +1,32 @@ +package com.twitter.util.validation.constraints; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +import com.twitter.util.validation.internal.validators.StateConstraintValidator; + +import jakarta.validation.Constraint; +import jakarta.validation.Payload; + +@Target(ElementType.PARAMETER) +@Retention(RetentionPolicy.RUNTIME) +@Constraint(validatedBy = { StateConstraintValidator.class }) +public @interface StateConstraint { + + /** Every constraint annotation must define a message element of type String. */ + String message() default "Please register with state CA"; + + /** + * Every constraint annotation must define a groups element that specifies the processing groups + * with which the constraint declaration is associated. + */ + Class[] groups() default {}; + + /** + * Constraint annotations must define a payload element that specifies the payload with which + * the constraint declaration is associated. * + */ + Class[] payload() default {}; +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraintPayload.scala b/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraintPayload.scala new file mode 100644 index 0000000000..85d1a5d03d --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/constraints/StateConstraintPayload.scala @@ -0,0 +1,8 @@ +package com.twitter.util.validation.constraints + +import jakarta.validation.Payload + +case class StateConstraintPayload( + invalid: String, + expected: String) + extends Payload diff --git a/util-validator/src/test/scala/com/twitter/util/validation/conversions/PathOpsTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/conversions/PathOpsTest.scala new file mode 100644 index 0000000000..15a949d5eb --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/conversions/PathOpsTest.scala @@ -0,0 +1,58 @@ +package com.twitter.util.validation.conversions + +import com.twitter.util.validation.conversions.PathOps._ +import jakarta.validation.Path +import org.hibernate.validator.internal.engine.path.PathImpl +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class PathOpsTest extends AnyFunSuite with Matchers { + + test("PathOps#root") { + val path: Path = PathImpl.createRootPath() + + val leafNode = path.getLeafNode + leafNode should not be (null) + leafNode.toString should equal("") + } + + test("PathOps#without leaf node") { + val path = PathImpl.createRootPath() + + // only the root exists, so this fails + intercept[IndexOutOfBoundsException] { + PathImpl.createCopyWithoutLeafNode(path) + } + + // add nodes + path.addBeanNode() + path.addPropertyNode("foo") + + val pathWithLeaf: Path = PathImpl.createCopy(path) + val leafNode = pathWithLeaf.getLeafNode + leafNode should not be (null) + leafNode.toString should equal("foo") + } + + test("PathOps#propertyNode") { + val path = PathImpl.createRootPath() + path.addPropertyNode("leaf") + + val propertyPath: Path = PathImpl.createCopy(path) + + val leafNode = propertyPath.getLeafNode + leafNode should not be (null) + leafNode.toString should equal("leaf") + } + + test("PathOps#fromString") { + val path: Path = PathImpl.createPathFromString("foo.bar.baz") + + val leafNode = path.getLeafNode + leafNode should not be (null) + leafNode.toString should equal("baz") + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/executable/ScalaExecutableValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/executable/ScalaExecutableValidatorTest.scala new file mode 100644 index 0000000000..137c0b78f5 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/executable/ScalaExecutableValidatorTest.scala @@ -0,0 +1,891 @@ +package com.twitter.util.validation.executable + +import com.twitter.util.validation.AssertViolationTest.WithViolation +import com.twitter.util.validation.ScalaValidatorTest.DefaultAddress +import com.twitter.util.validation.caseclasses.{ + Car, + CarWithPassengerCount, + Customer, + NestedUser, + Page, + Path, + PathNotEmpty, + Person, + RandoMixin, + RentalStation, + RentalStationMixin, + UnicodeNameCaseClass, + UsersRequest +} +import com.twitter.util.validation.cfg.ConstraintMapping +import com.twitter.util.validation.constraints.ValidPassengerCountReturnValue +import com.twitter.util.validation.engine.ConstraintViolationHelper +import com.twitter.util.validation.internal.validators.ValidPassengerCountReturnValueConstraintValidator +import com.twitter.util.validation.{AssertViolationTest, CarMake, ScalaValidator} +import jakarta.validation.{ConstraintViolation, UnexpectedTypeException} +import java.time.LocalDate + +class ScalaExecutableValidatorTest extends AssertViolationTest { + + override protected val validator: ScalaValidator = + ScalaValidator.builder + .withConstraintMappings( + Set( + ConstraintMapping( + classOf[ValidPassengerCountReturnValue], + classOf[ValidPassengerCountReturnValueConstraintValidator] + ) + ) + ) + .validator + + private[this] val executableValidator: ScalaExecutableValidator = validator.forExecutables + + test("ScalaExecutableValidator#validateMethods 1") { + intercept[IllegalArgumentException] { + executableValidator.validateMethods[Car](null) + } + } + + test("ScalaExecutableValidator#validateMethods 2") { + val owner = Person(id = "", name = "A. Einstein", address = DefaultAddress) + val car = Car( + id = 1234, + make = CarMake.Volkswagen, + model = "Beetle", + year = 1970, + owners = Seq(owner), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyEnd = Some(LocalDate.now.plusYears(3)), + passengers = Seq( + Person(id = "1001", name = "R. Franklin", address = DefaultAddress) + ) + ) + + val violations = executableValidator.validateMethods(car) + assertViolations( + violations, + Seq( + WithViolation("validateId", "id may not be even", car, classOf[Car], car), + WithViolation( + "warrantyTimeValid.warrantyEnd", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car), + WithViolation( + "warrantyTimeValid.warrantyStart", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car) + ) + ) + } + + test("ScalaExecutableValidator#validateMethods 3") { + val owner = Person(id = "", name = "A. Einstein", address = DefaultAddress) + val car = Car( + id = 1234, + make = CarMake.Volkswagen, + model = "Beetle", + year = 1970, + owners = Seq(owner), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyEnd = Some(LocalDate.now.plusYears(3)), + passengers = Seq( + Person(id = "1001", name = "R. Franklin", address = DefaultAddress) + ) + ) + + val methods = validator.describeMethods(classOf[Car]) + val violations = executableValidator.validateMethods(methods, car) + assertViolations( + violations, + Seq( + WithViolation("validateId", "id may not be even", car, classOf[Car], car), + WithViolation( + "warrantyTimeValid.warrantyEnd", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car), + WithViolation( + "warrantyTimeValid.warrantyStart", + "both warrantyStart and warrantyEnd are required for a valid range", + car, + classOf[Car], + car) + ) + ) + } + + test("ScalaExecutableValidator#validateMethods 4") { + val nestedUser = NestedUser( + "abcd1234", + Person(id = "1234abcd", name = "A. Einstein", address = DefaultAddress), + "Other", + "" + ) + val methods = validator.describeMethods(classOf[NestedUser]) + val violations = executableValidator.validateMethods(methods, nestedUser) + assertViolations( + violations, + Seq( + WithViolation( + "jobCheck.job", + "cannot be empty", + nestedUser, + classOf[NestedUser], + nestedUser) + ) + ) + } + + test("ScalaExecutableValidator#validateMethod") { + val owner = Person(id = "", name = "A. Einstein", address = DefaultAddress) + val car = Car( + id = 1234, + make = CarMake.Volkswagen, + model = "Beetle", + year = 1970, + owners = Seq(owner), + licensePlate = "CA123", + numDoors = 2, + manual = true, + ownershipStart = LocalDate.now, + ownershipEnd = LocalDate.now.plusYears(10), + warrantyEnd = Some(LocalDate.now.plusYears(3)), + passengers = Seq( + Person(id = "1001", name = "R. Franklin", address = DefaultAddress) + ) + ) + + val method = classOf[Car].getDeclaredMethod("validateId") + + val violations = executableValidator.validateMethod(car, method) + assertViolations( + violations, + Seq( + WithViolation("validateId", "id may not be even", car, classOf[Car], car) + ) + ) + } + + test("ScalaExecutableValidator#validateParameters 1") { + val customer = Customer("Jane", "Smith") + val rentalStation = RentalStation("Hertz") + + val method = + classOf[RentalStation].getMethod( + "rentCar", + classOf[Customer], + classOf[LocalDate], + classOf[Int]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array(customer, LocalDate.now().minusDays(1), Integer.valueOf(5))) + violations.size should equal(1) + val invalidRentalStartDateViolation = violations.head + invalidRentalStartDateViolation.getMessage should equal("must be a future date") + invalidRentalStartDateViolation.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation.getRootBean should equal(rentalStation) + invalidRentalStartDateViolation.getLeafBean should equal(rentalStation) + + val invalidRentalStartDateAndNullCustomerViolations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array(null, LocalDate.now(), Integer.valueOf(5))) + invalidRentalStartDateAndNullCustomerViolations.size should equal(2) + val sortedInvalidRentalStartDateAndNullCustomerViolationsViolations = + ConstraintViolationHelper.sortViolations[RentalStation]( + invalidRentalStartDateAndNullCustomerViolations) + val nullCustomerViolation = sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.head + nullCustomerViolation.getMessage should equal("must not be null") + nullCustomerViolation.getPropertyPath.toString should equal("rentCar.customer") + nullCustomerViolation.getRootBeanClass should equal(classOf[RentalStation]) + nullCustomerViolation.getRootBean should equal(rentalStation) + nullCustomerViolation.getLeafBean should equal(rentalStation) + + val invalidRentalStartDateViolation2 = + sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.last + invalidRentalStartDateViolation2.getMessage should equal("must be a future date") + invalidRentalStartDateViolation2.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation2.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation2.getRootBean should equal(rentalStation) + invalidRentalStartDateViolation2.getLeafBean should equal(rentalStation) + } + + test("ScalaExecutableValidator#validateParameters 2") { + val rentalStation = RentalStation("Hertz") + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array(Seq.empty[Customer]) + ) + violations.size should equal(2) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should equal(rentalStation) + violation1.getInvalidValue should equal(Seq.empty[Customer]) + + val violation2 = sortedViolations.last + violation2.getMessage should equal("size must be between 1 and 2147483647") + violation2.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation2.getRootBeanClass should equal(classOf[RentalStation]) + violation2.getRootBean should equal(rentalStation) + violation2.getInvalidValue should equal(Seq.empty[Customer]) + } + + test("ScalaExecutableValidator#validateParameters 3") { + // validation should cascade to customer parameter which is missing first name + val rentalStation = RentalStation("Hertz") + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array( + Seq( + Customer("", "Ride") + )) + ) + violations.size should equal(1) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers[0].first") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should equal(rentalStation) + violation1.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateParameters 4") { + // validation should handle encoded parameter name + val rentalStation = RentalStation("Hertz") + val method = + classOf[RentalStation].getMethod("listCars", classOf[String]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateParameters( + rentalStation, + method, + Array("") + ) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("listCars.airport-code") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should equal(rentalStation) + violation.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateReturnValue") { + val rentalStation = RentalStation("Hertz") + + val method = + classOf[RentalStation].getMethod("getCustomers") + val returnValue = Seq.empty[Customer] + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateReturnValue( + rentalStation, + method, + returnValue + ) + violations.size should equal(2) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("getCustomers.") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should equal(rentalStation) + violation1.getInvalidValue should equal(Seq.empty[Customer]) + + val violation2 = sortedViolations.last + violation2.getMessage should equal("size must be between 1 and 2147483647") + violation2.getPropertyPath.toString should equal("getCustomers.") + violation2.getRootBeanClass should equal(classOf[RentalStation]) + violation2.getRootBean should equal(rentalStation) + violation2.getInvalidValue should equal(Seq.empty[Customer]) + } + + test("ScalaExecutableValidator#validateConstructorParameters 1") { + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateConstructorParameters( + constructor, + Array(null) + ) + violations.size should equal(1) + + val violation = violations.head + violation.getMessage should equal("must not be null") + violation.getPropertyPath.toString should equal("RentalStation.name") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should be(null) + } + + test("ScalaExecutableValidator#validateConstructorParameters 2") { + val constructor = classOf[UnicodeNameCaseClass].getConstructor(classOf[Int], classOf[String]) + + val violations: Set[ConstraintViolation[UnicodeNameCaseClass]] = + executableValidator.validateConstructorParameters( + constructor, + Array(5, "") + ) + violations.size should equal(2) + + val sortedViolations: Seq[ConstraintViolation[UnicodeNameCaseClass]] = + ConstraintViolationHelper.sortViolations(violations) + + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("UnicodeNameCaseClass.name") + violation1.getRootBeanClass should equal(classOf[UnicodeNameCaseClass]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal("") + + val violation2 = sortedViolations(1) + violation2.getMessage should equal("must be less than or equal to 3") + violation2.getPropertyPath.toString should equal("UnicodeNameCaseClass.winning-id") + violation2.getRootBeanClass should equal(classOf[UnicodeNameCaseClass]) + violation2.getRootBean should be(null) + violation2.getInvalidValue should equal(5) + } + + test("ScalaExecutableValidator#validateConstructorParameters 3") { + val constructor = classOf[Page[String]].getConstructor( + classOf[List[String]], + classOf[Int], + classOf[Option[Long]], + classOf[Option[Long]]) + + val violations: Set[ConstraintViolation[Page[String]]] = + executableValidator.validateConstructorParameters( + constructor, + Array( + List("foo", "bar"), + 0.asInstanceOf[AnyRef], + Some(1), + Some(2) + ) + ) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should be("Page.pageSize") + violation.getMessage should be("must be greater than or equal to 1") + violation.getRootBeanClass should equal(classOf[Page[Person]]) + violation.getRootBean should be(null) + violation.getInvalidValue should equal(0) + } + + test("ScalaExecutableValidator#validateMethodParameters 1") { + // validation should handle encoded parameter name + val method = + classOf[RentalStation].getMethod("listCars", classOf[String]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array("") + ) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("listCars.airport-code") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateMethodParameters 2") { + // validation should cascade to customer parameter which is missing first name + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array( + Seq( + Customer("", "Ride") + )) + ) + violations.size should equal(1) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers[0].first") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateMethodParameters 3") { + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array(Seq.empty[Customer]) + ) + violations.size should equal(2) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal(Seq.empty[Customer]) + + val violation2 = sortedViolations.last + violation2.getMessage should equal("size must be between 1 and 2147483647") + violation2.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation2.getRootBeanClass should equal(classOf[RentalStation]) + violation2.getRootBean should be(null) + violation2.getInvalidValue should equal(Seq.empty[Customer]) + } + + test("ScalaExecutableValidator#validateMethodParameters 4") { + val customer = Customer("Jane", "Smith") + + val method = + classOf[RentalStation].getMethod( + "rentCar", + classOf[Customer], + classOf[LocalDate], + classOf[Int]) + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array(customer, LocalDate.now().minusDays(1), Integer.valueOf(5))) + violations.size should equal(1) + val invalidRentalStartDateViolation = violations.head + invalidRentalStartDateViolation.getMessage should equal("must be a future date") + invalidRentalStartDateViolation.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation.getRootBean should be(null) + invalidRentalStartDateViolation.getLeafBean should be(null) + + val invalidRentalStartDateAndNullCustomerViolations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array(null, LocalDate.now(), Integer.valueOf(5))) + invalidRentalStartDateAndNullCustomerViolations.size should equal(2) + val sortedInvalidRentalStartDateAndNullCustomerViolationsViolations: Seq[ + ConstraintViolation[RentalStation] + ] = ConstraintViolationHelper.sortViolations(invalidRentalStartDateAndNullCustomerViolations) + + val nullCustomerViolation = sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.head + nullCustomerViolation.getMessage should equal("must not be null") + nullCustomerViolation.getPropertyPath.toString should equal("rentCar.customer") + nullCustomerViolation.getRootBeanClass should equal(classOf[RentalStation]) + nullCustomerViolation.getRootBean should be(null) + nullCustomerViolation.getLeafBean should be(null) + + val invalidRentalStartDateViolation2 = + sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.last + invalidRentalStartDateViolation2.getMessage should equal("must be a future date") + invalidRentalStartDateViolation2.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation2.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation2.getRootBean should be(null) + invalidRentalStartDateViolation2.getLeafBean should be(null) + } + + test("ScalaExecutableValidator#validateExecutableParameters 1") { + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + val descriptor = validator.describeExecutable(constructor, None) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array(null), + constructor.getParameters.map(_.getName) + ) + violations.size should equal(1) + + val violation = violations.head + violation.getMessage should equal("must not be null") + violation.getPropertyPath.toString should equal("RentalStation.name") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should be(null) + } + + test("ScalaExecutableValidator#validateExecutableParameters 2") { + val constructor = classOf[UnicodeNameCaseClass].getConstructor(classOf[Int], classOf[String]) + val descriptor = validator.describeExecutable(constructor, None) + + val violations: Set[ConstraintViolation[UnicodeNameCaseClass]] = + executableValidator.validateExecutableParameters( + descriptor, + Array(5, ""), + constructor.getParameters.map(_.getName) + ) + violations.size should equal(2) + + val sortedViolations: Seq[ConstraintViolation[UnicodeNameCaseClass]] = + ConstraintViolationHelper.sortViolations(violations) + + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("UnicodeNameCaseClass.name") + violation1.getRootBeanClass should equal(classOf[UnicodeNameCaseClass]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal("") + + val violation2 = sortedViolations(1) + violation2.getMessage should equal("must be less than or equal to 3") + violation2.getPropertyPath.toString should equal("UnicodeNameCaseClass.winning-id") + violation2.getRootBeanClass should equal(classOf[UnicodeNameCaseClass]) + violation2.getRootBean should be(null) + violation2.getInvalidValue should equal(5) + } + + test("ScalaExecutableValidator#validateExecutableParameters 3") { + val constructor = classOf[Page[String]].getConstructor( + classOf[List[String]], + classOf[Int], + classOf[Option[Long]], + classOf[Option[Long]]) + val descriptor = validator.describeExecutable(constructor, None) + + val violations: Set[ConstraintViolation[Page[String]]] = + executableValidator.validateExecutableParameters( + descriptor, + Array( + List("foo", "bar"), + 0.asInstanceOf[AnyRef], + Some(1), + Some(2) + ), + constructor.getParameters.map(_.getName) + ) + violations.size should equal(1) + val violation = violations.head + violation.getPropertyPath.toString should be("Page.pageSize") + violation.getMessage should be("must be greater than or equal to 1") + violation.getRootBeanClass should equal(classOf[Page[Person]]) + violation.getRootBean should be(null) + violation.getInvalidValue should equal(0) + } + + test("ScalaExecutableValidator#validateExecutableParameters 4") { + // validation should handle encoded parameter name + val method = + classOf[RentalStation].getMethod("listCars", classOf[String]) + val descriptor = validator.describeExecutable(method, None) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array(""), + method.getParameters.map(_.getName) + ) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("must not be empty") + violation.getPropertyPath.toString should equal("listCars.airport-code") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateExecutableParameters 5") { + // validation should cascade to customer parameter which is missing first name + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val descriptor = validator.describeExecutable(method, None) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array( + Seq( + Customer("", "Ride") + ) + ), + method.getParameters.map(_.getName) + ) + violations.size should equal(1) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers[0].first") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal("") + } + + test("ScalaExecutableValidator#validateExecutableParameters 6") { + val method = + classOf[RentalStation].getMethod("updateCustomerRecords", classOf[Seq[Customer]]) + val descriptor = validator.describeExecutable(method, None) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array(Seq.empty[Customer]), + method.getParameters.map(_.getName) + ) + violations.size should equal(2) + val sortedViolations: Seq[ConstraintViolation[RentalStation]] = + ConstraintViolationHelper.sortViolations(violations) + val violation1 = sortedViolations.head + violation1.getMessage should equal("must not be empty") + violation1.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation1.getRootBeanClass should equal(classOf[RentalStation]) + violation1.getRootBean should be(null) + violation1.getInvalidValue should equal(Seq.empty[Customer]) + + val violation2 = sortedViolations.last + violation2.getMessage should equal("size must be between 1 and 2147483647") + violation2.getPropertyPath.toString should equal("updateCustomerRecords.customers") + violation2.getRootBeanClass should equal(classOf[RentalStation]) + violation2.getRootBean should be(null) + violation2.getInvalidValue should equal(Seq.empty[Customer]) + } + + test("ScalaExecutableValidator#validateExecutableParameters 7") { + val customer = Customer("Jane", "Smith") + + val method = + classOf[RentalStation].getMethod( + "rentCar", + classOf[Customer], + classOf[LocalDate], + classOf[Int]) + val descriptor = validator.describeExecutable(method, None) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array(customer, LocalDate.now().minusDays(1), Integer.valueOf(5)), + method.getParameters.map(_.getName) + ) + violations.size should equal(1) + val invalidRentalStartDateViolation = violations.head + invalidRentalStartDateViolation.getMessage should equal("must be a future date") + invalidRentalStartDateViolation.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation.getRootBean should be(null) + invalidRentalStartDateViolation.getLeafBean should be(null) + + val invalidRentalStartDateAndNullCustomerViolations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateMethodParameters( + method, + Array(null, LocalDate.now(), Integer.valueOf(5))) + invalidRentalStartDateAndNullCustomerViolations.size should equal(2) + val sortedInvalidRentalStartDateAndNullCustomerViolationsViolations: Seq[ + ConstraintViolation[RentalStation] + ] = ConstraintViolationHelper.sortViolations(invalidRentalStartDateAndNullCustomerViolations) + + val nullCustomerViolation = sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.head + nullCustomerViolation.getMessage should equal("must not be null") + nullCustomerViolation.getPropertyPath.toString should equal("rentCar.customer") + nullCustomerViolation.getRootBeanClass should equal(classOf[RentalStation]) + nullCustomerViolation.getRootBean should be(null) + nullCustomerViolation.getLeafBean should be(null) + + val invalidRentalStartDateViolation2 = + sortedInvalidRentalStartDateAndNullCustomerViolationsViolations.last + invalidRentalStartDateViolation2.getMessage should equal("must be a future date") + invalidRentalStartDateViolation2.getPropertyPath.toString should equal("rentCar.start") + invalidRentalStartDateViolation2.getRootBeanClass should equal(classOf[RentalStation]) + invalidRentalStartDateViolation2.getRootBean should be(null) + invalidRentalStartDateViolation2.getLeafBean should be(null) + } + + test("ScalaExecutableValidator#validateExecutableParameters 8") { + // test mixin class support + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + val descriptor = validator.describeExecutable(constructor, Some(classOf[RentalStationMixin])) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array("Hertz"), + constructor.getParameters.map(_.getName) + ) + violations.size should equal(1) + + // this violates the @Min(10) constraint from RentalStationMixin + val violation = violations.head + violation.getMessage should equal("must be greater than or equal to 10") + violation.getPropertyPath.toString should equal("RentalStation.name") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should be("Hertz") + } + + test("ScalaExecutableValidator#validateExecutableParameters 9") { + // test mixin class support and replacement of parameter names + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + val descriptor = validator.describeExecutable(constructor, Some(classOf[RentalStationMixin])) + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array("Hertz"), + Array("id") + ) + violations.size should equal(1) + + // this violates the @Min(10) constraint from RentalStationMixin + val violation = violations.head + violation.getMessage should equal("must be greater than or equal to 10") + violation.getPropertyPath.toString should equal("RentalStation.id") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should be(null) + violation.getInvalidValue should be("Hertz") + } + + test("ScalaExecutableValidator#validateExecutableParameters 10") { + // test mixin class support with mixin that doesn't define the field + val constructor = classOf[RentalStation].getConstructor(classOf[String]) + val descriptor = + validator.describeExecutable( + constructor, + Some(classOf[RandoMixin]) + ) // doesn't define any fields so no additional constraints + + val violations: Set[ConstraintViolation[RentalStation]] = + executableValidator.validateExecutableParameters( + descriptor, + Array("Hertz"), + Array("id") + ) + violations.isEmpty should be(true) + } + + test("ScalaExecutableValidator#validateExecutableParameters 11") { + val pathValue = Path.apply("foo") + val constructor = classOf[PathNotEmpty].getConstructor(classOf[Path], classOf[String]) + val descriptor = validator.describeExecutable(constructor, None) + + // no validator registered `@NotEmpty` for Path type, should be returned as a violation not an exception + val e = intercept[UnexpectedTypeException] { + executableValidator.validateExecutableParameters( + descriptor, + Array(pathValue, "12345"), + constructor.getParameters.map(_.getName) + ) + } + e.getMessage should equal( + s"No validator could be found for constraint 'interface jakarta.validation.constraints.NotEmpty' validating type 'com.twitter.util.validation.caseclasses$$Path'. Check configuration for 'PathNotEmpty.path'") + } + + test("ScalaExecutableValidator#validateExecutableParameters 12") { + val constructor = classOf[UsersRequest].getConstructor( + classOf[Int], + classOf[Option[LocalDate]], + classOf[Boolean]) + val descriptor = validator.describeExecutable(constructor, None) + + executableValidator + .validateExecutableParameters( + descriptor, + Array( + 10, + Some(LocalDate.parse("2013-01-01")), + true + ), + constructor.getParameters.map(_.getName) + ).isEmpty should be(true) + + val violations: Set[ConstraintViolation[UsersRequest]] = + executableValidator + .validateExecutableParameters( + descriptor, + Array( + 200, + Some(LocalDate.parse("2013-01-01")), + true + ), + constructor.getParameters.map(_.getName) + ) + violations.size should equal(1) + + val violation = violations.head + violation.getMessage should equal("must be less than or equal to 100") + violation.getPropertyPath.toString should equal("UsersRequest.max") + violation.getRootBeanClass should equal(classOf[UsersRequest]) + violation.getRootBean should be(null) + violation.getInvalidValue should be(200) + } + + test("ScalaExecutableValidator#validateConstructorReturnValue") { + val car = CarWithPassengerCount( + Seq( + Person("abcd1234", "J Doe", DefaultAddress), + Person("abcd1235", "K Doe", DefaultAddress), + Person("abcd1236", "L Doe", DefaultAddress), + Person("abcd1237", "M Doe", DefaultAddress), + Person("abcd1238", "N Doe", DefaultAddress) + ) + ) + val constructor = classOf[CarWithPassengerCount].getConstructor(classOf[Seq[Person]]) + + val violations: Set[ConstraintViolation[CarWithPassengerCount]] = + executableValidator.validateConstructorReturnValue( + constructor, + car + ) + violations.size should equal(1) + + val violation = violations.head + violation.getMessage should equal("number of passenger(s) is not valid") + violation.getPropertyPath.toString should equal("CarWithPassengerCount.") + violation.getRootBeanClass should equal(classOf[CarWithPassengerCount]) + violation.getRootBean should equal(null) + violation.getInvalidValue should be(car) + } + + test("ScalaExecutableValidator#cross-parameter") { + val rentalStation = RentalStation("Hertz") + val method = + classOf[RentalStation].getMethod("reserve", classOf[LocalDate], classOf[LocalDate]) + + val start = LocalDate.now.plusDays(1) + val end = LocalDate.now + val violations = executableValidator.validateParameters( + rentalStation, + method, + Array(start, end) + ) + violations.size should equal(1) + val violation = violations.head + violation.getMessage should equal("start is not before end") + violation.getPropertyPath.toString should equal("reserve.") + violation.getRootBeanClass should equal(classOf[RentalStation]) + violation.getRootBean should equal(rentalStation) + violation.getLeafBean should equal(rentalStation) + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConsistentDateParametersValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConsistentDateParametersValidator.scala new file mode 100644 index 0000000000..51a6b8fcba --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConsistentDateParametersValidator.scala @@ -0,0 +1,21 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.constraints.ConsistentDateParameters +import jakarta.validation.constraintvalidation.{SupportedValidationTarget, ValidationTarget} +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import java.time.LocalDate + +@SupportedValidationTarget(Array(ValidationTarget.PARAMETERS)) +class ConsistentDateParametersValidator + extends ConstraintValidator[ConsistentDateParameters, Array[AnyRef]] { + override def isValid( + value: Array[AnyRef], + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + if (value.length != 2) throw new IllegalArgumentException("Unexpected method signature") + else if (value(0) == null || value(1) == null) true + else if (!value(0).isInstanceOf[LocalDate] || !value(1).isInstanceOf[LocalDate]) + throw new IllegalArgumentException("Unexpected method signature") + else value(0).asInstanceOf[LocalDate].isBefore(value(1).asInstanceOf[LocalDate]) + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConstraintValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConstraintValidatorTest.scala new file mode 100644 index 0000000000..1d7fbc3b09 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ConstraintValidatorTest.scala @@ -0,0 +1,29 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.ScalaValidator +import jakarta.validation.ConstraintViolation +import java.lang.annotation.Annotation +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner +import org.scalatestplus.scalacheck.ScalaCheckDrivenPropertyChecks +import scala.reflect.runtime.universe._ + +@RunWith(classOf[JUnitRunner]) +abstract class ConstraintValidatorTest + extends AnyFunSuite + with Matchers + with ScalaCheckDrivenPropertyChecks { + + protected def validator: ScalaValidator + + protected def testFieldName: String + + protected def validate[A <: Annotation, T: TypeTag]( + clazz: Class[_], + paramName: String, + value: Any + ): Set[ConstraintViolation[T]] = + validator.validateValue(clazz.asInstanceOf[Class[T]], paramName, value) +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidatorTest.scala new file mode 100644 index 0000000000..a5b16ce029 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ISO3166CountryCodeConstraintValidatorTest.scala @@ -0,0 +1,107 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.reflect.{Types => ReflectTypes} +import com.twitter.util.validation.ScalaValidator +import com.twitter.util.validation.caseclasses.{ + CountryCodeArrayExample, + CountryCodeExample, + CountryCodeInvalidTypeExample, + CountryCodeSeqExample +} +import com.twitter.util.validation.constraints.CountryCode +import jakarta.validation.ConstraintViolation +import java.util.Locale +import org.scalacheck.Gen +import org.scalacheck.Shrink.shrinkAny +import scala.reflect.runtime.universe._ + +class ISO3166CountryCodeConstraintValidatorTest extends ConstraintValidatorTest { + protected val validator: ScalaValidator = ScalaValidator() + protected val testFieldName: String = "countryCode" + + private val countryCodes = Locale.getISOCountries.filter(_.nonEmpty).toSeq + + test("pass validation for valid country code") { + countryCodes.foreach { value => + validate[CountryCodeExample](value).isEmpty should be(true) + } + } + + test("fail validation for invalid country code") { + forAll(genFakeCountryCode) { value => + val violations = validate[CountryCodeExample](value) + violations.size should equal(1) + violations.head.getMessage should equal(s"$value not a valid country code") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + test("pass validation for valid country codes in seq") { + val passValue = Gen.nonEmptyContainerOf[Seq, String](Gen.oneOf(countryCodes)) + + forAll(passValue) { value => + validate[CountryCodeSeqExample](value).isEmpty should be(true) + } + } + + test("fail validation for empty seq") { + val emptyValue = Seq.empty + val violations = validate[CountryCodeSeqExample](emptyValue) + violations.size should equal(1) + violations.head.getMessage should equal(" not a valid country code") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + + test("fail validation for invalid country codes in seq") { + val failValue = Gen.nonEmptyContainerOf[Seq, String](genFakeCountryCode) + + forAll(failValue) { value => + val violations = validate[CountryCodeSeqExample](value) + violations.size should equal(1) + violations.head.getMessage should equal(s"${mkString(value)} not a valid country code") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + test("pass validation for valid country codes in array") { + val passValue = Gen.nonEmptyContainerOf[Array, String](Gen.oneOf(countryCodes)) + + forAll(passValue) { value => + validate[CountryCodeArrayExample](value).isEmpty should be(true) + } + } + + test("fail validation for invalid country codes in array") { + val failValue = Gen.nonEmptyContainerOf[Array, String](genFakeCountryCode) + + forAll(failValue) { value => + val violations = validate[CountryCodeArrayExample](value) + violations.size should equal(1) + violations.head.getMessage should equal(s"${mkString(value)} not a valid country code") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + test("fail validation for invalid country code type") { + val failValue = Gen.choose[Int](0, 100) + + forAll(failValue) { value => + val violations = validate[CountryCodeInvalidTypeExample](value) + violations.size should equal(1) + violations.head.getMessage should equal(s"${value.toString} not a valid country code") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + //generate random uppercase string for fake country code + private def genFakeCountryCode: Gen[String] = { + Gen + .nonEmptyContainerOf[Seq, Char](Gen.alphaUpperChar) + .map(_.mkString) + .filter(!countryCodes.contains(_)) + } + + private def validate[T: TypeTag](value: Any): Set[ConstraintViolation[T]] = { + super.validate[CountryCode, T](ReflectTypes.runtimeClass[T], testFieldName, value) + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/InvalidConstraintValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/InvalidConstraintValidator.scala new file mode 100644 index 0000000000..2975dba8c3 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/InvalidConstraintValidator.scala @@ -0,0 +1,12 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.constraints.InvalidConstraint +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import scala.util.control.NoStackTrace + +class InvalidConstraintValidator extends ConstraintValidator[InvalidConstraint, Any] { + override def isValid( + obj: Any, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = throw new RuntimeException("validator foo error") with NoStackTrace +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/NotEmptyPathConstraintValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/NotEmptyPathConstraintValidator.scala new file mode 100644 index 0000000000..fc9ad037ec --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/NotEmptyPathConstraintValidator.scala @@ -0,0 +1,25 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.caseclasses.Path +import jakarta.validation.constraints.NotEmpty +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + +class NotEmptyPathConstraintValidator extends ConstraintValidator[NotEmpty, Path] { + override def isValid( + obj: Path, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = obj match { + case Path.Empty => false + case _ => true + } +} + +class NotEmptyAnyConstraintValidator extends ConstraintValidator[NotEmpty, Any] { + override def isValid( + obj: Any, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = obj match { + case Path.Empty => false + case _ => true + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidatorTest.scala new file mode 100644 index 0000000000..141891b7af --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/OneOfConstraintValidatorTest.scala @@ -0,0 +1,78 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.reflect.{Types => ReflectTypes} +import com.twitter.util.validation.ScalaValidator +import com.twitter.util.validation.caseclasses._ +import com.twitter.util.validation.constraints.OneOf +import jakarta.validation.ConstraintViolation +import org.scalacheck.Gen +import org.scalacheck.Shrink.shrinkAny +import scala.reflect.runtime.universe._ + +class OneOfConstraintValidatorTest extends ConstraintValidatorTest { + + protected val validator: ScalaValidator = ScalaValidator() + protected val testFieldName: String = "enumValue" + + val oneOfValues: Set[String] = Set("a", "B", "c") + + test("pass validation for single value") { + oneOfValues.foreach { value => + validate[OneOfExample](value).isEmpty should be(true) + } + } + + test("fail validation for single value") { + val failValue = Gen.alphaStr.filter(!oneOfValues.contains(_)) + + forAll(failValue) { value => + val violations = validate[OneOfExample](value) + violations.size should equal(1) + violations.head.getMessage.endsWith("not one of [a, B, c]") should be(true) + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + test("pass validation for seq") { + val passValue = Gen.nonEmptyContainerOf[Seq, String](Gen.oneOf(oneOfValues.toSeq)) + + forAll(passValue) { value => + validate[OneOfSeqExample](value) + } + } + + test("fail validation for empty seq") { + val emptySeq = Seq.empty + + val violations = validate[OneOfSeqExample](emptySeq) + violations.size should equal(1) + violations.head.getMessage.endsWith("not one of [a, B, c]") should be(true) + violations.head.getPropertyPath.toString should equal(testFieldName) + } + + test("fail validation for invalid value in seq") { + val failValue = + Gen.nonEmptyContainerOf[Seq, String](Gen.alphaStr.filter(!oneOfValues.contains(_))) + + forAll(failValue) { value => + val violations = validate[OneOfSeqExample](value) + violations.size should equal(1) + violations.head.getMessage.endsWith("not one of [a, B, c]") should be(true) + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + test("fail validation for invalid type") { + val failValue = Gen.choose(0, 100) + + forAll(failValue) { value => + val violations = validate[OneOfInvalidTypeExample](value) + violations.size should equal(1) + violations.head.getMessage.endsWith("not one of [a, B, c]") should be(true) + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + private def validate[T: TypeTag](value: Any): Set[ConstraintViolation[T]] = + super.validate[OneOf, T](ReflectTypes.runtimeClass[T], testFieldName, value) +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/PackageObjectTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/PackageObjectTest.scala new file mode 100644 index 0000000000..2bd27c1f07 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/PackageObjectTest.scala @@ -0,0 +1,69 @@ +package com.twitter.util.validation.internal.validators + +import org.junit.runner.RunWith +import org.scalatest.funsuite.AnyFunSuite +import org.scalatest.matchers.should.Matchers +import org.scalatestplus.junit.JUnitRunner + +@RunWith(classOf[JUnitRunner]) +class PackageObjectTest extends AnyFunSuite with Matchers { + + test("mkString with more than 1 entry - Seq") { + mkString(Seq(1, 2)) should equal("[1, 2]") + } + + test("mkString with map entries") { + val m: Map[String, String] = Map("foo" -> "1", "bar" -> "2", "baz" -> "3") + mkString(m) should equal("[(foo,1), (bar,2), (baz,3)]") + } + + test("mkString with option") { + val o = Some("foo") + mkString(o) should equal("foo") + } + + test("mkString with more than 1 entry - Seq some empty values") { + mkString(Seq(" ", " 33")) should equal("33") + mkString(Seq(" ", " 33", " ")) should equal("33") + mkString(Seq(" ", " ", " ", " 33", " ")) should equal("33") + } + + test("mkString with more than 1 entry - Array") { + mkString(Array(1, 2)) should equal("[1, 2]") + } + + test("mkString with 1 entry - Seq") { + mkString(Seq(1)) should equal("1") + mkString(Seq(-1)) should equal("-1") + mkString(Seq("!!")) should equal("!!") + } + + test("mkString with 1 entry - Array") { + mkString(Array(1)) should equal("1") + } + + test("mkString with no entries - Seq") { + mkString(Seq(" ", " ")) should equal("") + mkString(Seq.empty[Int]) should equal("") + mkString(Seq.empty[String]) should equal("") + } + + test("mkString with no entries - Array") { + mkString(Array(" ", " ")) should equal("") + mkString(Array.empty[Int]) should equal("") + mkString(Array.empty[String]) should equal("") + } + + test("mkString with no entries - Map") { + mkString(Map("" -> "")) should equal("(,)") + mkString(Map.empty[String, Int]) should equal("") + mkString(Map.empty) should equal("") + } + + test("mkString with no entries - Option") { + mkString(None) should equal("") + mkString(Option(null)) should equal("") + mkString(Option("")) should equal("") + mkString(Some("")) should equal("") + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/StateConstraintValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/StateConstraintValidator.scala new file mode 100644 index 0000000000..91665b67b6 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/StateConstraintValidator.scala @@ -0,0 +1,23 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.constraints.{StateConstraint, StateConstraintPayload} +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} +import org.hibernate.validator.constraintvalidation.HibernateConstraintValidatorContext + +class StateConstraintValidator extends ConstraintValidator[StateConstraint, String] { + + override def isValid( + obj: String, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = { + val valid = obj.equalsIgnoreCase("CA") + + if (!valid) { + val hibernateContext: HibernateConstraintValidatorContext = + constraintValidatorContext.unwrap(classOf[HibernateConstraintValidatorContext]) + hibernateContext.withDynamicPayload(StateConstraintPayload(obj, "CA")) + } + + valid + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidatorTest.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidatorTest.scala new file mode 100644 index 0000000000..343358a50e --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/UUIDConstraintValidatorTest.scala @@ -0,0 +1,38 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.reflect.{Types => ReflectTypes} +import com.twitter.util.validation.ScalaValidator +import com.twitter.util.validation.caseclasses._ +import com.twitter.util.validation.constraints.UUID +import jakarta.validation.ConstraintViolation +import org.scalacheck.Gen +import scala.reflect.runtime.universe._ + +class UUIDConstraintValidatorTest extends ConstraintValidatorTest { + + protected val validator: ScalaValidator = ScalaValidator() + protected val testFieldName: String = "uuid" + + test("pass validation for valid uuid") { + val passValue = Gen.uuid + + forAll(passValue) { value => + validate[UUIDExample](value.toString).isEmpty should be(true) + } + } + + test("fail validation for invalid uuid") { + val passValue = Gen.alphaStr + + forAll(passValue) { value => + val violations = validate[UUIDExample](value) + violations.size should equal(1) + violations.head.getMessage should equal("must be a valid UUID") + violations.head.getPropertyPath.toString should equal(testFieldName) + } + } + + private def validate[T: TypeTag](value: String): Set[ConstraintViolation[T]] = { + super.validate[UUID, T](ReflectTypes.runtimeClass[T], testFieldName, value) + } +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountConstraintValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountConstraintValidator.scala new file mode 100644 index 0000000000..d3fcc55ae9 --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountConstraintValidator.scala @@ -0,0 +1,21 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.caseclasses.Car +import com.twitter.util.validation.constraints.ValidPassengerCount +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + +class ValidPassengerCountConstraintValidator extends ConstraintValidator[ValidPassengerCount, Car] { + + @volatile private[this] var maxPassengers: Long = _ + + override def initialize(constraintAnnotation: ValidPassengerCount): Unit = { + this.maxPassengers = constraintAnnotation.max() + } + + override def isValid( + car: Car, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = + if (car == null) true + else car.passengers.nonEmpty && car.passengers.size <= this.maxPassengers +} diff --git a/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountReturnValueConstraintValidator.scala b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountReturnValueConstraintValidator.scala new file mode 100644 index 0000000000..99bbd77d0a --- /dev/null +++ b/util-validator/src/test/scala/com/twitter/util/validation/internal/validators/ValidPassengerCountReturnValueConstraintValidator.scala @@ -0,0 +1,22 @@ +package com.twitter.util.validation.internal.validators + +import com.twitter.util.validation.caseclasses.CarWithPassengerCount +import com.twitter.util.validation.constraints.ValidPassengerCountReturnValue +import jakarta.validation.constraintvalidation.{SupportedValidationTarget, ValidationTarget} +import jakarta.validation.{ConstraintValidator, ConstraintValidatorContext} + +@SupportedValidationTarget(Array(ValidationTarget.ANNOTATED_ELEMENT)) +class ValidPassengerCountReturnValueConstraintValidator + extends ConstraintValidator[ValidPassengerCountReturnValue, CarWithPassengerCount] { + + @volatile private[this] var maxPassengers: Long = _ + + override def initialize(constraintAnnotation: ValidPassengerCountReturnValue): Unit = { + this.maxPassengers = constraintAnnotation.max() + } + + override def isValid( + obj: CarWithPassengerCount, + constraintValidatorContext: ConstraintValidatorContext + ): Boolean = obj.passengers.nonEmpty && obj.passengers.size <= this.maxPassengers +} diff --git a/util-zk-test/BUILD b/util-zk-test/BUILD index 88fa6a0060..36b2472230 100644 --- a/util-zk-test/BUILD +++ b/util-zk-test/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-zk-test/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-zk-test/src/test/java", "util/util-zk-test/src/test/scala", diff --git a/util-zk-test/PROJECT b/util-zk-test/PROJECT index aaa4d791f9..1aaeccf787 100644 --- a/util-zk-test/PROJECT +++ b/util-zk-test/PROJECT @@ -1,6 +1,4 @@ owners: - - coord-team:ldap - - banderson - - koliver - - mnakamura - - roanta + - coord-review:ldap + - jillianc + - mbezoyan diff --git a/util-zk-test/src/main/scala/BUILD b/util-zk-test/src/main/scala/BUILD index 3e8171168f..95e7cbf554 100644 --- a/util-zk-test/src/main/scala/BUILD +++ b/util-zk-test/src/main/scala/BUILD @@ -1,15 +1,6 @@ -scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, - provides = scala_artifact( - org = "com.twitter", - name = "util-zk-test", - repo = artifactory, - ), +target( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/apache/zookeeper:zookeeper-server", - ], - exports = [ - "3rdparty/jvm/org/apache/zookeeper:zookeeper-server", + "util/util-zk-test/src/main/scala/com/twitter/zk", ], ) diff --git a/util-zk-test/src/main/scala/com/twitter/zk/BUILD b/util-zk-test/src/main/scala/com/twitter/zk/BUILD new file mode 100644 index 0000000000..2ce870403d --- /dev/null +++ b/util-zk-test/src/main/scala/com/twitter/zk/BUILD @@ -0,0 +1,17 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-zk-test", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-server", + ], + exports = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-server", + ], +) diff --git a/util-zk-test/src/test/java/BUILD b/util-zk-test/src/test/java/BUILD index e21fe2662c..07d960c355 100644 --- a/util-zk-test/src/test/java/BUILD +++ b/util-zk-test/src/test/java/BUILD @@ -1,8 +1,15 @@ junit_tests( - sources = rglobs("*.java"), - fatal_warnings = True, + sources = ["**/*.java"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", - "util/util-zk-test", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-zk-test/src/main/scala/com/twitter/zk", ], ) diff --git a/util-zk-test/src/test/scala/BUILD b/util-zk-test/src/test/scala/BUILD index ab2be0a286..cf92253e76 100644 --- a/util-zk-test/src/test/scala/BUILD +++ b/util-zk-test/src/test/scala/BUILD @@ -1,12 +1,18 @@ junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, + sources = ["**/*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], dependencies = [ "3rdparty/jvm/junit", "3rdparty/jvm/org/apache/zookeeper:zookeeper-server", - "3rdparty/jvm/org/mockito:mockito-all", "3rdparty/jvm/org/scalatest", - "util/util-core", - "util/util-zk-test", + "3rdparty/jvm/org/scalatestplus:junit", + "util/util-core/src/main/scala/com/twitter/io", + "util/util-zk-test/src/main/scala/com/twitter/zk", ], ) diff --git a/util-zk-test/src/test/scala/com/twitter/zk/ServerCnxnFactoryTest.scala b/util-zk-test/src/test/scala/com/twitter/zk/ServerCnxnFactoryTest.scala index 4aa27f6b81..a6f9991aa9 100644 --- a/util-zk-test/src/test/scala/com/twitter/zk/ServerCnxnFactoryTest.scala +++ b/util-zk-test/src/test/scala/com/twitter/zk/ServerCnxnFactoryTest.scala @@ -4,12 +4,10 @@ import com.twitter.io.TempDirectory import java.io.File import java.net.InetAddress import org.apache.zookeeper.server.ZooKeeperServer -import org.junit.runner.RunWith -import org.scalatest.{BeforeAndAfter, FunSuite} -import org.scalatest.junit.JUnitRunner +import org.scalatest.BeforeAndAfter +import org.scalatest.funsuite.AnyFunSuite -@RunWith(classOf[JUnitRunner]) -class ServerCnxnFactoryTest extends FunSuite with BeforeAndAfter { +class ServerCnxnFactoryTest extends AnyFunSuite with BeforeAndAfter { val addr = InetAddress.getLocalHost var testServer: ZooKeeperServer = null diff --git a/util-zk/BUILD b/util-zk/BUILD index e9907b6656..b251269408 100644 --- a/util-zk/BUILD +++ b/util-zk/BUILD @@ -1,11 +1,13 @@ target( + tags = ["bazel-compatible"], dependencies = [ "util/util-zk/src/main/scala", ], ) -target( +test_suite( name = "tests", + tags = ["bazel-compatible"], dependencies = [ "util/util-zk/src/test/scala", ], diff --git a/util-zk/GROUPS b/util-zk/GROUPS deleted file mode 100644 index 82a0574ba7..0000000000 --- a/util-zk/GROUPS +++ /dev/null @@ -1,3 +0,0 @@ -coord-team -csl -observability diff --git a/util-zk/OWNERS b/util-zk/OWNERS deleted file mode 100644 index a3a086672a..0000000000 --- a/util-zk/OWNERS +++ /dev/null @@ -1,6 +0,0 @@ -banderson -dschobel -koliver -mnakamura -roanta -stuhood diff --git a/util-zk/PROJECT b/util-zk/PROJECT new file mode 100644 index 0000000000..08d46e03da --- /dev/null +++ b/util-zk/PROJECT @@ -0,0 +1,9 @@ +# See go/project-docs. Generated using ./pants run src/python/twitter/phabricator:owners_to_project -- --dirs=a/owned/path,b/owned/path + +owners: + - jillianc + - mbezoyan +watchers: + - coord-team@twitter.com + - csl-team:ldap + - observability@twitter.com diff --git a/util-zk/src/main/scala/BUILD b/util-zk/src/main/scala/BUILD index bc1401a136..ac89a3d878 100644 --- a/util-zk/src/main/scala/BUILD +++ b/util-zk/src/main/scala/BUILD @@ -1,18 +1,20 @@ scala_library( - sources = rglobs("*.scala"), - fatal_warnings = True, + # This target exists only to act as an aggregator jar. + sources = ["PantsWorkaroundZk.scala"], + platform = "java8", provides = scala_artifact( org = "com.twitter", name = "util-zk", repo = artifactory, ), + strict_deps = True, + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", - "util/util-core/src/main/scala", - "util/util-logging/src/main/scala", + "util/util-zk/src/main/scala/com/twitter/zk", + "util/util-zk/src/main/scala/com/twitter/zk/coordination", ], exports = [ - "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", - "util/util-core/src/main/scala", + "util/util-zk/src/main/scala/com/twitter/zk", + "util/util-zk/src/main/scala/com/twitter/zk/coordination", ], ) diff --git a/util-zk/src/main/scala/PantsWorkaroundZk.scala b/util-zk/src/main/scala/PantsWorkaroundZk.scala new file mode 100644 index 0000000000..202482a88b --- /dev/null +++ b/util-zk/src/main/scala/PantsWorkaroundZk.scala @@ -0,0 +1,4 @@ +// workaround for https://github.com/pantsbuild/pants/issues/6678 +private final class PantsWorkaroundZk { + throw new IllegalStateException("workaround for https://github.com/pantsbuild/pants/issues/6678") +} diff --git a/util-zk/src/main/scala/com/twitter/zk/AsyncCallbackPromise.scala b/util-zk/src/main/scala/com/twitter/zk/AsyncCallbackPromise.scala index 7ce12f3fab..e7a70debfc 100644 --- a/util-zk/src/main/scala/com/twitter/zk/AsyncCallbackPromise.scala +++ b/util-zk/src/main/scala/com/twitter/zk/AsyncCallbackPromise.scala @@ -2,7 +2,7 @@ package com.twitter.zk import java.util.{List => JList} -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ import org.apache.zookeeper.data.Stat import org.apache.zookeeper.{AsyncCallback, KeeperException} @@ -23,13 +23,13 @@ trait AsyncCallbackPromise[T] extends Promise[T] { class StringCallbackPromise extends AsyncCallbackPromise[String] with AsyncCallback.StringCallback { def processResult(rc: Int, path: String, ctx: AnyRef, name: String): Unit = { - process(rc, path) { name } + process(rc, path)(name) } } class UnitCallbackPromise extends AsyncCallbackPromise[Unit] with AsyncCallback.VoidCallback { def processResult(rc: Int, path: String, ctx: AnyRef): Unit = { - process(rc, path) { Unit } + process(rc, path)(()) } } @@ -46,7 +46,13 @@ class ExistsCallbackPromise(znode: ZNode) class ChildrenCallbackPromise(znode: ZNode) extends AsyncCallbackPromise[ZNode.Children] with AsyncCallback.Children2Callback { - def processResult(rc: Int, path: String, ctx: AnyRef, children: JList[String], stat: Stat): Unit = { + def processResult( + rc: Int, + path: String, + ctx: AnyRef, + children: JList[String], + stat: Stat + ): Unit = { process(rc, path) { znode(stat, children.asScala.toSeq) } diff --git a/util-zk/src/main/scala/com/twitter/zk/AuthInfo.scala b/util-zk/src/main/scala/com/twitter/zk/AuthInfo.scala index a94d6ba41d..ab6488f87e 100644 --- a/util-zk/src/main/scala/com/twitter/zk/AuthInfo.scala +++ b/util-zk/src/main/scala/com/twitter/zk/AuthInfo.scala @@ -3,8 +3,12 @@ package com.twitter.zk case class AuthInfo(mode: String, data: Array[Byte]) object AuthInfo { - def digest(user: String, secret: String) = { - val authString = "%s:%s".format(user, secret).getBytes("UTF-8") //DigestAuthenticationProvider.generateDigest("%s:%s".format(user, secret)).getBytes("UTF-8") + def digest(user: String, secret: String): AuthInfo = { + val authString = + "%s:%s" + .format(user, secret).getBytes( + "UTF-8" + ) //DigestAuthenticationProvider.generateDigest("%s:%s".format(user, secret)).getBytes("UTF-8") AuthInfo("digest", authString) } } diff --git a/util-zk/src/main/scala/com/twitter/zk/BUILD b/util-zk/src/main/scala/com/twitter/zk/BUILD new file mode 100644 index 0000000000..0041ab88b4 --- /dev/null +++ b/util-zk/src/main/scala/com/twitter/zk/BUILD @@ -0,0 +1,24 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-zk-zk", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-logging/src/main/scala/com/twitter/logging", + ], + exports = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-core/src/main/scala/com/twitter/conversions", + ], +) diff --git a/util-zk/src/main/scala/com/twitter/zk/Connector.scala b/util-zk/src/main/scala/com/twitter/zk/Connector.scala index 8170bd1b6b..33b9432df5 100644 --- a/util-zk/src/main/scala/com/twitter/zk/Connector.scala +++ b/util-zk/src/main/scala/com/twitter/zk/Connector.scala @@ -14,7 +14,7 @@ trait Connector { import Connector.EventHandler val name = "zk-connector" - protected lazy val log = Logger.get(name) + protected lazy val log: Logger = Logger.get(name) private[this] val listeners = new AtomicReference[List[EventHandler]](Nil) @@ -60,10 +60,10 @@ object Connector { */ case class RoundRobin(connectors: Connector*) extends Connector { require(connectors.length > 0) - override val name = "round-robin-zk-connector:%d".format(connectors.length) + override val name: String = "round-robin-zk-connector:%d".format(connectors.length) private[this] var index = 0 - protected[this] def nextConnector() = { + protected[this] def nextConnector(): Connector = { val i = synchronized { if (index == Int.MaxValue) { index = 0 diff --git a/util-zk/src/main/scala/com/twitter/zk/Event.scala b/util-zk/src/main/scala/com/twitter/zk/Event.scala index b206cbe4d7..72c70b1c1f 100644 --- a/util-zk/src/main/scala/com/twitter/zk/Event.scala +++ b/util-zk/src/main/scala/com/twitter/zk/Event.scala @@ -11,7 +11,8 @@ import com.twitter.util.{Promise, Return} */ object Event { - def apply(t: EventType, s: KeeperState, p: Option[String]) = new WatchedEvent(t, s, p.orNull) + def apply(t: EventType, s: KeeperState, p: Option[String]): WatchedEvent = + new WatchedEvent(t, s, p.orNull) def unapply(event: WatchedEvent): Option[(EventType, KeeperState, Option[String])] = { Some((event.getType, event.getState, Option { event.getPath })) @@ -19,10 +20,10 @@ object Event { } sealed trait StateEvent { - val eventType = EventType.None + val eventType: EventType = EventType.None val state: KeeperState - def apply() = Event(eventType, state, None) - def unapply(event: WatchedEvent) = event match { + def apply(): WatchedEvent = Event(eventType, state, None) + def unapply(event: WatchedEvent): Boolean = event match { case Event(t, s, _) => (t == eventType && s == state) case _ => false } @@ -30,52 +31,51 @@ sealed trait StateEvent { object StateEvent { object AuthFailed extends StateEvent { - val state = KeeperState.AuthFailed + val state: KeeperState = KeeperState.AuthFailed } object Connected extends StateEvent { - val state = KeeperState.SyncConnected + val state: KeeperState = KeeperState.SyncConnected } object Disconnected extends StateEvent { - val state = KeeperState.Disconnected + val state: KeeperState = KeeperState.Disconnected } object Expired extends StateEvent { - val state = KeeperState.Expired + val state: KeeperState = KeeperState.Expired } object ConnectedReadOnly extends StateEvent { - val state = KeeperState.ConnectedReadOnly + val state: KeeperState = KeeperState.ConnectedReadOnly } object SaslAuthenticated extends StateEvent { - val state = KeeperState.SaslAuthenticated + val state: KeeperState = KeeperState.SaslAuthenticated } def apply(w: WatchedEvent): StateEvent = { - w.getState match { + val state = w.getState + state match { case KeeperState.AuthFailed => AuthFailed case KeeperState.SyncConnected => Connected case KeeperState.Disconnected => Disconnected case KeeperState.Expired => Expired case KeeperState.ConnectedReadOnly => ConnectedReadOnly case KeeperState.SaslAuthenticated => SaslAuthenticated - case KeeperState.Unknown => - throw new IllegalArgumentException("Can't convert deprecated state to StateEvent: Unknown") - case KeeperState.NoSyncConnected => - throw new IllegalArgumentException( - "Can't convert deprecated state to StateEvent: NoSyncConnected" - ) + case _ => + throw new IllegalArgumentException(s"Can't convert deprecated state to StateEvent: $state") + //NoSyncConnected and Unknown are depricated in zk 3.x, and should be + //expected to be removed in zk 4.x } } } sealed trait NodeEvent { - val state = KeeperState.SyncConnected + val state: KeeperState = KeeperState.SyncConnected val eventType: EventType - def apply(path: String) = Event(eventType, state, Some(path)) - def unapply(event: WatchedEvent) = event match { + def apply(path: String): WatchedEvent = Event(eventType, state, Some(path)) + def unapply(event: WatchedEvent): Option[String] = event match { case Event(t, _, somePath) if (t == eventType) => somePath case _ => None } @@ -83,19 +83,19 @@ sealed trait NodeEvent { object NodeEvent { object Created extends NodeEvent { - val eventType = EventType.NodeCreated + val eventType: EventType = EventType.NodeCreated } object ChildrenChanged extends NodeEvent { - val eventType = EventType.NodeChildrenChanged + val eventType: EventType = EventType.NodeChildrenChanged } object DataChanged extends NodeEvent { - val eventType = EventType.NodeDataChanged + val eventType: EventType = EventType.NodeDataChanged } object Deleted extends NodeEvent { - val eventType = EventType.NodeDeleted + val eventType: EventType = EventType.NodeDeleted } } diff --git a/util-zk/src/main/scala/com/twitter/zk/LiftableFuture.scala b/util-zk/src/main/scala/com/twitter/zk/LiftableFuture.scala index b23867dce2..1390aa5951 100644 --- a/util-zk/src/main/scala/com/twitter/zk/LiftableFuture.scala +++ b/util-zk/src/main/scala/com/twitter/zk/LiftableFuture.scala @@ -3,6 +3,7 @@ package com.twitter.zk import com.twitter.util.{Future, Return, Throw} import org.apache.zookeeper.KeeperException import scala.language.implicitConversions +import com.twitter.util.Try protected[zk] object LiftableFuture { implicit def liftableFuture[T](f: Future[T]): LiftableFuture[T] = new LiftableFuture(f) @@ -15,14 +16,18 @@ protected[zk] object LiftableFuture { protected[zk] class LiftableFuture[T](f: Future[T]) { /** Lift a value to a Return. */ - def liftSuccess = f map { Return(_) } + def liftSuccess: Future[Return[T]] = f map { Return(_) } /** Lift all errors to a Throw */ - def liftFailure = liftSuccess handle { case e => Throw(e) } + def liftFailure: Future[Try[T]] = liftSuccess handle { case e => Throw(e) } /** Lift all KeeperExceptions to a Throw */ - def liftKeeperException = liftSuccess handle { case e: KeeperException => Throw(e) } + def liftKeeperException: Future[Try[T]] = liftSuccess handle { + case e: KeeperException => Throw(e) + } /** Lift failures when a watch would have been successfully installed */ - def liftNoNode = liftSuccess handle { case e: KeeperException.NoNodeException => Throw(e) } + def liftNoNode: Future[Try[T]] = liftSuccess handle { + case e: KeeperException.NoNodeException => Throw(e) + } } diff --git a/util-zk/src/main/scala/com/twitter/zk/NativeConnector.scala b/util-zk/src/main/scala/com/twitter/zk/NativeConnector.scala index e9f4123788..ef6755a404 100644 --- a/util-zk/src/main/scala/com/twitter/zk/NativeConnector.scala +++ b/util-zk/src/main/scala/com/twitter/zk/NativeConnector.scala @@ -13,12 +13,12 @@ case class NativeConnector( connectTimeout: Option[Duration], sessionTimeout: Duration, timer: Timer, - authenticate: Option[AuthInfo] = None -) extends Connector + authenticate: Option[AuthInfo] = None) + extends Connector with Serialized { override val name = "native-zk-connector" - protected[this] def mkConnection = { + protected[this] def mkConnection: NativeConnector.Connection = { new NativeConnector.Connection( connectString, connectTimeout, @@ -48,20 +48,16 @@ case class NativeConnector( */ def apply(): Future[ZooKeeper] = serialized { - connection getOrElse { + connection.getOrElse { val c = mkConnection - c.sessionEvents foreach { event => - sessionBroker.send(event()).sync() - } + c.sessionEvents foreach { event => sessionBroker.send(event()).sync() } connection = Some(c) c - } apply () + }.apply }.flatten .rescue { case e: NativeConnector.ConnectTimeoutException => - release() flatMap { _ => - Future.exception(e) - } + release() flatMap { _ => Future.exception(e) } } /** @@ -79,13 +75,20 @@ case class NativeConnector( } object NativeConnector { - def apply(connectString: String, sessionTimeout: Duration)( + def apply( + connectString: String, + sessionTimeout: Duration + )( implicit timer: Timer ): NativeConnector = { NativeConnector(connectString, None, sessionTimeout, timer) } - def apply(connectString: String, connectTimeout: Duration, sessionTimeout: Duration)( + def apply( + connectString: String, + connectTimeout: Duration, + sessionTimeout: Duration + )( implicit timer: Timer ): NativeConnector = { NativeConnector(connectString, Some(connectTimeout), sessionTimeout, timer) @@ -96,7 +99,9 @@ object NativeConnector { connectTimeout: Duration, sessionTimeout: Duration, authenticate: Option[AuthInfo] - )(implicit timer: Timer): NativeConnector = { + )( + implicit timer: Timer + ): NativeConnector = { NativeConnector(connectString, Some(connectTimeout), sessionTimeout, timer, authenticate) } @@ -115,25 +120,26 @@ object NativeConnector { sessionTimeout: Duration, authenticate: Option[AuthInfo], timer: Timer, - log: Logger - ) { + log: Logger) { @volatile protected[this] var zookeeper: Option[ZooKeeper] = None - protected[this] val connectPromise = new Promise[ZooKeeper] + protected[this] val connectPromise: Promise[ZooKeeper] = new Promise[ZooKeeper] /** A ZooKeeper handle that will error if connectTimeout is specified and exceeded. */ - lazy val connected: Future[ZooKeeper] = connectTimeout.map { timeout => - connectPromise.within(timer, timeout).rescue { - case _: TimeoutException => - Future.exception(ConnectTimeoutException(connectString, timeout)) + lazy val connected: Future[ZooKeeper] = connectTimeout + .map { timeout => + connectPromise.within(timer, timeout).rescue { + case _: TimeoutException => + Future.exception(ConnectTimeoutException(connectString, timeout)) + } } - }.getOrElse(connectPromise) + .getOrElse(connectPromise) - protected[this] val releasePromise = new Promise[Unit] + protected[this] val releasePromise: Promise[Unit] = new Promise[Unit] val released: Future[Unit] = releasePromise - protected[this] val sessionBroker = new EventBroker + protected[this] val sessionBroker: EventBroker = new EventBroker /** * Publish session events, but intercept Connected events and use them to satisfy the pending @@ -142,7 +148,8 @@ object NativeConnector { val sessionEvents: Offer[StateEvent] = { val broker = new Broker[StateEvent] def loop(): Unit = { - sessionBroker.recv.sync() + sessionBroker.recv + .sync() .map { StateEvent(_) } .respond { case Return(StateEvent.Connected) => @@ -174,7 +181,7 @@ object NativeConnector { connected } - protected[this] def mkZooKeeper = { + protected[this] def mkZooKeeper: ZooKeeper = { new ZooKeeper(connectString, sessionTimeout.inMillis.toInt, sessionBroker) } @@ -183,7 +190,7 @@ object NativeConnector { log.debug("release") zk.close() zookeeper = None - releasePromise.setValue(Unit) + releasePromise.setValue(()) } } } diff --git a/util-zk/src/main/scala/com/twitter/zk/RetryPolicy.scala b/util-zk/src/main/scala/com/twitter/zk/RetryPolicy.scala index ca8eba31ad..fa2e77e1a8 100644 --- a/util-zk/src/main/scala/com/twitter/zk/RetryPolicy.scala +++ b/util-zk/src/main/scala/com/twitter/zk/RetryPolicy.scala @@ -2,7 +2,7 @@ package com.twitter.zk import org.apache.zookeeper.KeeperException -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Duration, Future, Timer} /** Pluggable retry strategy. */ @@ -12,7 +12,7 @@ trait RetryPolicy { /** Matcher for connection-related KeeperExceptions. */ object KeeperConnectionException { - def unapply(e: KeeperException) = e match { + def unapply(e: KeeperException): Option[KeeperException] = e match { case e: KeeperException.ConnectionLossException => Some(e) case e: KeeperException.SessionExpiredException => Some(e) case e: KeeperException.SessionMovedException => Some(e) @@ -46,7 +46,8 @@ object RetryPolicy { base: Duration, factor: Double = 2.0, maximum: Duration = 30.seconds - )(implicit timer: Timer) + )( + implicit timer: Timer) extends RetryPolicy { require(base > 0.seconds) require(factor >= 1) @@ -68,6 +69,6 @@ object RetryPolicy { /** A single try */ object None extends RetryPolicy { - def apply[T](op: => Future[T]) = op + def apply[T](op: => Future[T]): Future[T] = op } } diff --git a/util-zk/src/main/scala/com/twitter/zk/ZNode.scala b/util-zk/src/main/scala/com/twitter/zk/ZNode.scala index 50b64eb9b3..5b5fe3fb58 100644 --- a/util-zk/src/main/scala/com/twitter/zk/ZNode.scala +++ b/util-zk/src/main/scala/com/twitter/zk/ZNode.scala @@ -1,7 +1,7 @@ package com.twitter.zk -import scala.collection.JavaConverters._ import scala.collection.{Seq, Set} +import scala.jdk.CollectionConverters._ import org.apache.zookeeper.common.PathUtils import org.apache.zookeeper.data.{ACL, Stat} @@ -9,6 +9,7 @@ import org.apache.zookeeper.{CreateMode, KeeperException, WatchedEvent} import com.twitter.concurrent.{Broker, Offer} import com.twitter.util.{Future, Return, Throw, Try} +import com.twitter.logging.Logger /** * A handle to a ZNode attached to a ZkClient @@ -19,14 +20,14 @@ trait ZNode { val path: String protected[zk] val zkClient: ZkClient - protected[this] lazy val log = zkClient.log + protected[this] lazy val log: Logger = zkClient.log - override def hashCode = path.hashCode + override def hashCode: Int = path.hashCode - override def toString = "ZNode(%s)".format(path) + override def toString: String = "ZNode(%s)".format(path) /** ZNodes are equal if they share a path. */ - override def equals(other: Any) = other match { + override def equals(other: Any): Boolean = other match { case z @ ZNode(_) => (z.hashCode == hashCode) case _ => false } @@ -36,7 +37,7 @@ trait ZNode { */ /** Return the ZkClient associated with this node. */ - def client = zkClient + def client: ZkClient = zkClient /** Get a child node. */ def apply(child: String): ZNode = ZNode(zkClient, childPath(child)) @@ -90,9 +91,7 @@ trait ZNode { zkClient.retrying { zk => val result = new StringCallbackPromise zk.create(creatingPath, data, acls.asJava, mode, result, null) - result map { newPath => - zkClient(newPath) - } + result map { newPath => zkClient(newPath) } } } @@ -100,9 +99,7 @@ trait ZNode { def delete(version: Int = 0): Future[ZNode] = zkClient.retrying { zk => val result = new UnitCallbackPromise zk.delete(path, version, result, null) - result map { _ => - this - } + result map { _ => this } } /** Returns a Future that is satisfied with this ZNode with its metadata and data */ @@ -116,9 +113,7 @@ trait ZNode { def sync(): Future[ZNode] = zkClient.retrying { zk => val result = new UnitCallbackPromise zk.sync(path, result, null) - result map { _ => - this - } + result map { _ => this } } /** Provides access to this node's children. */ @@ -214,13 +209,14 @@ trait ZNode { /** Pipe events from a subtree's monitor to this broker. */ def pipeSubTreeUpdates(next: Offer[ZNode.TreeUpdate]): Unit = { - next.sync().flatMap(broker ! _).onSuccess { _ => - pipeSubTreeUpdates(next) - } + next.sync().flatMap(broker ! _).onSuccess { _ => pipeSubTreeUpdates(next) } } /** Monitor a watch on this node. */ - def monitorWatch(watch: Future[ZNode.Watch[ZNode.Children]], knownChildren: Set[ZNode]): Unit = { + def monitorWatch( + watch: Future[ZNode.Watch[ZNode.Children]], + knownChildren: Set[ZNode] + ): Unit = { log.debug("monitoring %s with %d known children", path, knownChildren.size) watch onFailure { e => // An error occurred and there's not really anything we can do about it. @@ -236,11 +232,9 @@ trait ZNode { removed = knownChildren -- children ) log.debug("updating %s with %d children", path, treeUpdate.added.size) - broker send (treeUpdate) sync () onSuccess { _ => + broker.send(treeUpdate).sync.onSuccess { _ => log.debug("updated %s with %d children", path, treeUpdate.added.size) - treeUpdate.added foreach { z => - pipeSubTreeUpdates(z.monitorTree()) - } + treeUpdate.added foreach { z => pipeSubTreeUpdates(z.monitorTree()) } eventUpdate onSuccess { event => log.debug("event received on %s: %s", path, event) } onSuccess { @@ -253,7 +247,7 @@ trait ZNode { // Tell the broker about the children we lost; otherwise, if there were no children, // this deletion should be reflected in a watch on the parent node, if one exists. if (knownChildren.size > 0) { - broker send (ZNode.TreeUpdate(this, removed = knownChildren)) sync () + broker.send(ZNode.TreeUpdate(this, removed = knownChildren)).sync } else { Future.Done } onSuccess { _ => @@ -272,7 +266,7 @@ trait ZNode { /** AuthFailed and Expired are unmonitorable. Everything else can be resumed. */ protected[this] object MonitorableEvent { - def unapply(event: WatchedEvent) = event match { + def unapply(event: WatchedEvent): Boolean = event match { case StateEvent.AuthFailed() => false case StateEvent.Expired() => false case _ => true @@ -286,25 +280,25 @@ trait ZNode { object ZNode { /** Build a ZNode */ - def apply(zk: ZkClient, _path: String) = new ZNode { + def apply(zk: ZkClient, _path: String): ZNode = new ZNode { PathUtils.validatePath(_path) protected[zk] val zkClient = zk val path = _path } /** matcher */ - def unapply(znode: ZNode) = Some(znode.path) + def unapply(znode: ZNode): Some[String] = Some(znode.path) /** A matcher for KeeperExceptions that have a non-null path. */ object Error { - def unapply(ke: KeeperException) = Option(ke.getPath) + def unapply(ke: KeeperException): Option[String] = Option(ke.getPath) } /** A ZNode with its Stat metadata. */ trait Exists extends ZNode { val stat: Stat - override def equals(other: Any) = other match { + override def equals(other: Any): Boolean = other match { case Exists(p, s) => (p == path && s == stat) case o => super.equals(o) } @@ -314,13 +308,13 @@ object ZNode { } object Exists { - def apply(znode: ZNode, _stat: Stat) = new Exists { + def apply(znode: ZNode, _stat: Stat): Exists = new Exists { val path = znode.path protected[zk] val zkClient = znode.zkClient val stat = _stat } def apply(znode: Exists): Exists = apply(znode, znode.stat) - def unapply(znode: Exists) = Some((znode.path, znode.stat)) + def unapply(znode: Exists): Some[(String, Stat)] = Some((znode.path, znode.stat)) } /** A ZNode with its Stat metadata and children znodes. */ @@ -328,7 +322,7 @@ object ZNode { val stat: Stat val children: Seq[ZNode] - override def equals(other: Any) = other match { + override def equals(other: Any): Boolean = other match { case Children(p, s, c) => (p == path && s == stat && c == children) case o => super.equals(o) } @@ -344,7 +338,7 @@ object ZNode { def apply(znode: ZNode, stat: Stat, children: Seq[String]): Children = { apply(Exists(znode, stat), children.map(znode.apply)) } - def unapply(z: Children) = Some((z.path, z.stat, z.children)) + def unapply(z: Children): Some[(String, Stat, Seq[ZNode])] = Some((z.path, z.stat, z.children)) } /** A ZNode with its Stat metadata and data. */ @@ -352,21 +346,22 @@ object ZNode { val stat: Stat val bytes: Array[Byte] - override def equals(other: Any) = other match { + override def equals(other: Any): Boolean = other match { case Data(p, s, b) => (p == path && s == stat && b == bytes) case o => super.equals(o) } } object Data { - def apply(znode: ZNode, _stat: Stat, _bytes: Array[Byte]) = new Data { + def apply(znode: ZNode, _stat: Stat, _bytes: Array[Byte]): Data = new Data { val path = znode.path protected[zk] val zkClient = znode.zkClient val stat = _stat val bytes = _bytes } def apply(znode: Exists, bytes: Array[Byte]): Data = apply(znode, znode.stat, bytes) - def unapply(znode: Data) = Some((znode.path, znode.stat, znode.bytes)) + def unapply(znode: Data): Some[(String, Stat, Array[Byte])] = + Some((znode.path, znode.stat, znode.bytes)) } case class Watch[T <: Exists](result: Try[T], update: Future[WatchedEvent]) { @@ -379,6 +374,5 @@ object ZNode { case class TreeUpdate( parent: ZNode, added: Set[ZNode] = Set.empty[ZNode], - removed: Set[ZNode] = Set.empty[ZNode] - ) + removed: Set[ZNode] = Set.empty[ZNode]) } diff --git a/util-zk/src/main/scala/com/twitter/zk/ZOp.scala b/util-zk/src/main/scala/com/twitter/zk/ZOp.scala index 889fa2acbe..3b9e6d0465 100644 --- a/util-zk/src/main/scala/com/twitter/zk/ZOp.scala +++ b/util-zk/src/main/scala/com/twitter/zk/ZOp.scala @@ -23,11 +23,7 @@ trait ZOp[T <: ZNode.Exists] { def setWatch(): Unit = { watch() onSuccess { case ZNode.Watch(result, update) => - broker ! result onSuccess { _ => - update onSuccess { _ => - setWatch() - } - } + broker ! result onSuccess { _ => update onSuccess { _ => setWatch() } } } } setWatch() diff --git a/util-zk/src/main/scala/com/twitter/zk/ZkClient.scala b/util-zk/src/main/scala/com/twitter/zk/ZkClient.scala index 7d86285ab1..ede3d426a9 100644 --- a/util-zk/src/main/scala/com/twitter/zk/ZkClient.scala +++ b/util-zk/src/main/scala/com/twitter/zk/ZkClient.scala @@ -1,6 +1,6 @@ package com.twitter.zk -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ import org.apache.zookeeper.ZooDefs.Ids.CREATOR_ALL_ACL import org.apache.zookeeper.data.ACL @@ -18,7 +18,7 @@ import com.twitter.util.{Duration, Future, Timer} */ trait ZkClient { val name = "zk.client" - protected[zk] val log = Logger.get(name) + protected[zk] val log: Logger = Logger.get(name) /** Build a ZNode handle using this ZkClient. */ def apply(path: String): ZNode = ZNode(this, path) @@ -42,7 +42,7 @@ trait ZkClient { /* Settings */ /** Default: Creator access */ - val acl: Seq[ACL] = CREATOR_ALL_ACL.asScala + val acl: Seq[ACL] = CREATOR_ALL_ACL.asScala.toSeq /** Build a ZkClient with another default creation ACL. */ def withAcl(a: Seq[ACL]): ZkClient = transform(_acl = a) @@ -71,7 +71,7 @@ trait ZkClient { _acl: Seq[ACL] = acl, _mode: CreateMode = mode, _retryPolicy: RetryPolicy = retryPolicy - ) = new ZkClient { + ): ZkClient = new ZkClient { val connector = _connector override val acl = _acl override val mode = _mode @@ -82,19 +82,27 @@ trait ZkClient { object ZkClient { /** Build a ZkClient with a provided Connector */ - def apply(_connector: Connector) = new ZkClient { + def apply(_connector: Connector): ZkClient = new ZkClient { protected[this] val connector = _connector } /** Build a ZkClient with a NativeConnector */ - def apply(connectString: String, connectTimeout: Option[Duration], sessionTimeout: Duration)( + def apply( + connectString: String, + connectTimeout: Option[Duration], + sessionTimeout: Duration + )( implicit timer: Timer ): ZkClient = { apply(NativeConnector(connectString, connectTimeout, sessionTimeout, timer)) } /** Build a ZkClient with a NativeConnector */ - def apply(connectString: String, connectTimeout: Duration, sessionTimeout: Duration)( + def apply( + connectString: String, + connectTimeout: Duration, + sessionTimeout: Duration + )( implicit timer: Timer ): ZkClient = { apply(connectString, Some(connectTimeout), sessionTimeout)(timer) diff --git a/util-zk/src/main/scala/com/twitter/zk/coordination/BUILD b/util-zk/src/main/scala/com/twitter/zk/coordination/BUILD new file mode 100644 index 0000000000..0c78a36744 --- /dev/null +++ b/util-zk/src/main/scala/com/twitter/zk/coordination/BUILD @@ -0,0 +1,23 @@ +scala_library( + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + provides = scala_artifact( + org = "com.twitter", + name = "util-zk-coordination", + repo = artifactory, + ), + strict_deps = True, + tags = ["bazel-compatible"], + dependencies = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-zk/src/main/scala/com/twitter/zk", + ], + exports = [ + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/concurrent", + "util/util-zk/src/main/scala/com/twitter/zk", + ], +) diff --git a/util-zk/src/main/scala/com/twitter/zk/coordination/ShardCoordinator.scala b/util-zk/src/main/scala/com/twitter/zk/coordination/ShardCoordinator.scala index 4ddc1b6d3d..7d02868971 100644 --- a/util-zk/src/main/scala/com/twitter/zk/coordination/ShardCoordinator.scala +++ b/util-zk/src/main/scala/com/twitter/zk/coordination/ShardCoordinator.scala @@ -48,8 +48,8 @@ class ShardCoordinator(zk: ZkClient, path: String, numShards: Int) { * client is finished performing work for the shard. * * @return A Future of ShardPermit that is satisfied when a shard slot becomes available. - * @throws SemaphoreError when the underlying semaphore throws an exception. - * @throws RejectedExecutionException when a shard cannot be acquired due to an unexpected + * @note Throws `SemaphoreError` when the underlying semaphore throws an exception. + * @note Throws `RejectedExecutionException` when a shard cannot be acquired due to an unexpected * state in the zookeeper tree. Assumed to only happen if some * zookeeper client clobbers the tree location for this * ShardCoordinator. @@ -57,9 +57,7 @@ class ShardCoordinator(zk: ZkClient, path: String, numShards: Int) { def acquire(): Future[ShardPermit] = { semaphore.acquire flatMap { permit => shardNodes() map { nodes => - nodes map { node => - shardIdOf(node.path) - } + nodes map { node => shardIdOf(node.path) } } map { ids => (0 until numShards) filterNot { ids contains _ } } flatMap { availableIds => @@ -84,9 +82,7 @@ class ShardCoordinator(zk: ZkClient, path: String, numShards: Int) { case err: LackOfConsensusException => Future.exception(SemaphoreError(err)) case err: PermitMismatchException => Future.exception(SemaphoreError(err)) case err: PermitNodeException => Future.exception(SemaphoreError(err)) - } onFailure { err => - permit.release() - } + } onFailure { err => permit.release() } } } @@ -100,9 +96,9 @@ class ShardCoordinator(zk: ZkClient, path: String, numShards: Int) { private[this] def shardNodes(): Future[Seq[ZNode]] = { zk(path).getChildren() map { zop => - zop.children filter { child => + (zop.children filter { child => child.path.startsWith(shardPathPrefix) - } sortBy (child => shardIdOf(child.path)) + } sortBy (child => shardIdOf(child.path))).toSeq } } @@ -121,7 +117,7 @@ sealed trait ShardPermit { case class Shard(id: Int, private val node: ZNode, private val permit: Permit) extends ShardPermit { - def release() = { + def release(): Unit = { node.delete() ensure { permit.release() } } } diff --git a/util-zk/src/main/scala/com/twitter/zk/coordination/ZkAsyncSemaphore.scala b/util-zk/src/main/scala/com/twitter/zk/coordination/ZkAsyncSemaphore.scala index 89f59aa1b9..9ce310cd42 100644 --- a/util-zk/src/main/scala/com/twitter/zk/coordination/ZkAsyncSemaphore.scala +++ b/util-zk/src/main/scala/com/twitter/zk/coordination/ZkAsyncSemaphore.scala @@ -35,7 +35,11 @@ import org.apache.zookeeper.{CreateMode, KeeperException} * } // handle { ... } * }}} */ -class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: Option[Int] = None) { +class ZkAsyncSemaphore( + zk: ZkClient, + path: String, + numPermits: Int, + maxWaiters: Option[Int] = None) { import ZkAsyncSemaphore._ require(numPermits > 0) require(maxWaiters.getOrElse(0) >= 0) @@ -47,8 +51,8 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: private[this] val waitq = new ConcurrentLinkedQueue[(Promise[ZkSemaphorePermit], ZNode)] private[this] class ZkSemaphorePermit(node: ZNode) extends Permit { - val zkPath = node.path - val sequenceNumber = sequenceNumberOf(zkPath) + val zkPath: String = node.path + val sequenceNumber: Int = sequenceNumberOf(zkPath) override def release(): Unit = { node.delete() } @@ -70,9 +74,7 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: val mySequenceNumber = sequenceNumberOf(permitNode.path) permitNodes() flatMap { permits => - val sequenceNumbers = permits map { child => - sequenceNumberOf(child.path) - } + val sequenceNumbers = permits map { child => sequenceNumberOf(child.path) } getConsensusNumPermits(permits) flatMap { consensusNumPermits => if (consensusNumPermits != numPermits) { throw ZkAsyncSemaphore.PermitMismatchException( @@ -120,9 +122,7 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: zk onSessionEvent { case StateEvent.Expired => rejectWaitQueue() case StateEvent.Connected => { - permitNodes() map { nodes => - checkWaiters(nodes) - } + permitNodes() map { nodes => checkWaiters(nodes) } monitorSemaphore(semaphoreNode) } } @@ -157,9 +157,9 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: * The sequence includes both nodes that have entered the semaphore as well as waiters. */ private[this] def permitNodes(): Future[Seq[ZNode]] = { - futureSemaphoreNode flatMap { semaphoreNode => + futureSemaphoreNode.flatMap { semaphoreNode => semaphoreNode.getChildren() map { zop => - zop.children filter { child => + zop.children.toSeq filter { child => child.path.startsWith(permitNodePathPrefix) } sortBy (child => sequenceNumberOf(child.path)) } @@ -175,11 +175,7 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: */ private[this] def monitorSemaphore(node: ZNode) = { val monitor = node.getChildren.monitor() - monitor foreach { tryChildren => - tryChildren map { zop => - checkWaiters(zop.children) - } - } + monitor foreach { tryChildren => tryChildren map { zop => checkWaiters(zop.children.toSeq) } } } /** @@ -201,9 +197,7 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: val permits = nodes filter { child => child.path.startsWith(permitNodePathPrefix) } sortBy (child => sequenceNumberOf(child.path)) - val ids = permits map { child => - sequenceNumberOf(child.path) - } + val ids = permits map { child => sequenceNumberOf(child.path) } val waitqIterator = waitq.iterator() while (waitqIterator.hasNext) { val (promise, permitNode) = waitqIterator.next() @@ -253,16 +247,14 @@ class ZkAsyncSemaphore(zk: ZkClient, path: String, numPermits: Int, maxWaiters: Future.collect(permits map numPermitsOf) map { purportedNumPermits => val groupedByNumPermits = purportedNumPermits filter { i => 0 < i - } groupBy { i => - i - } + } groupBy { i => i } val permitsToBelievers = groupedByNumPermits map { case (permits, believers) => (permits, believers.size) } val (numPermitsInMax, numBelieversOfMax) = permitsToBelievers.maxBy { case (_, believers) => believers } - val cardinalityOfMax = permitsToBelievers.values.filter({ _ == numBelieversOfMax }).size + val cardinalityOfMax = permitsToBelievers.values.count(_ == numBelieversOfMax) if (cardinalityOfMax == 1) { // Consensus diff --git a/util-zk/src/test/scala/BUILD b/util-zk/src/test/scala/BUILD index d81ef1f44e..9ef44003ba 100644 --- a/util-zk/src/test/scala/BUILD +++ b/util-zk/src/test/scala/BUILD @@ -1,14 +1,7 @@ -junit_tests( - sources = rglobs("*.scala"), - fatal_warnings = True, +test_suite( + tags = ["bazel-compatible"], dependencies = [ - "3rdparty/jvm/junit", - "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", - "3rdparty/jvm/org/jmock", - "3rdparty/jvm/org/mockito:mockito-all", - "3rdparty/jvm/org/scalatest", - "util/util-core/src/main/scala", - "util/util-logging/src/main/scala", - "util/util-zk/src/main/scala", + "util/util-zk/src/test/scala/com/twitter/zk", + "util/util-zk/src/test/scala/com/twitter/zk/coordination", ], ) diff --git a/util-zk/src/test/scala/com/twitter/zk/BUILD b/util-zk/src/test/scala/com/twitter/zk/BUILD new file mode 100644 index 0000000000..f705ddfd03 --- /dev/null +++ b/util-zk/src/test/scala/com/twitter/zk/BUILD @@ -0,0 +1,22 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-logging/src/main/scala/com/twitter/logging", + "util/util-zk/src/main/scala/com/twitter/zk", + ], +) diff --git a/util-zk/src/test/scala/com/twitter/zk/ConnectorTest.scala b/util-zk/src/test/scala/com/twitter/zk/ConnectorTest.scala index 6bd7074b54..5890f19ee0 100644 --- a/util-zk/src/test/scala/com/twitter/zk/ConnectorTest.scala +++ b/util-zk/src/test/scala/com/twitter/zk/ConnectorTest.scala @@ -1,15 +1,12 @@ package com.twitter.zk -import org.junit.runner.RunWith import org.mockito.Mockito._ -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar import com.twitter.util.Future +import org.scalatest.wordspec.AnyWordSpec -@RunWith(classOf[JUnitRunner]) -class ConnectorTest extends WordSpec with MockitoSugar { +class ConnectorTest extends AnyWordSpec with MockitoSugar { "Connector.RoundRobin" should { "require underlying connections" in { intercept[Exception] { @@ -26,9 +23,7 @@ class ConnectorTest extends WordSpec with MockitoSugar { connector } val nConnectors = 3 - val connectors = 1 to nConnectors map { _ => - mockConnector - } + val connectors = 1 to nConnectors map { _ => mockConnector } val connector = Connector.RoundRobin(connectors: _*) } @@ -36,30 +31,18 @@ class ConnectorTest extends WordSpec with MockitoSugar { val h = new ConnectorSpecHelper import h._ - connectors foreach { x => - assert(x.apply() == Future.never) - } - (1 to 2 * nConnectors) foreach { _ => - connector() - } - connectors foreach { c => - verify(c, times(3)).apply() - } + connectors foreach { x => assert(x.apply() == Future.never) } + (1 to 2 * nConnectors) foreach { _ => connector() } + connectors foreach { c => verify(c, times(3)).apply() } } "release" in { val h = new ConnectorSpecHelper import h._ - connectors foreach { x => - assert(x.release() == Future.never) - } - (1 to 2) foreach { _ => - connector.release() - } - connectors foreach { c => - verify(c, times(3)).release() - } + connectors foreach { x => assert(x.release() == Future.never) } + (1 to 2) foreach { _ => connector.release() } + connectors foreach { c => verify(c, times(3)).release() } } } } diff --git a/util-zk/src/test/scala/com/twitter/zk/ZNodeTest.scala b/util-zk/src/test/scala/com/twitter/zk/ZNodeTest.scala index 913c70c4b4..84c89beb75 100644 --- a/util-zk/src/test/scala/com/twitter/zk/ZNodeTest.scala +++ b/util-zk/src/test/scala/com/twitter/zk/ZNodeTest.scala @@ -3,13 +3,10 @@ package com.twitter.zk /** * @author ver@twitter.com */ -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar +import org.scalatest.wordspec.AnyWordSpec -@RunWith(classOf[JUnitRunner]) -class ZNodeTest extends WordSpec with MockitoSugar { +class ZNodeTest extends AnyWordSpec with MockitoSugar { "ZNode" should { class ZNodeSpecHelper { val zk = mock[ZkClient] @@ -33,9 +30,7 @@ class ZNodeTest extends WordSpec with MockitoSugar { val h = new ZNodeSpecHelper import h._ - val zs = (0 to 1) map { _ => - ZNode(zk, "/some/path") - } + val zs = (0 to 1) map { _ => ZNode(zk, "/some/path") } val table = Map(zs(0) -> true) assert(table.keys.toList.contains(zs(0))) assert(table.keys.toList.contains(zs(1))) diff --git a/util-zk/src/test/scala/com/twitter/zk/ZkClientTest.scala b/util-zk/src/test/scala/com/twitter/zk/ZkClientTest.scala index db4e949d6a..7915d044fc 100644 --- a/util-zk/src/test/scala/com/twitter/zk/ZkClientTest.scala +++ b/util-zk/src/test/scala/com/twitter/zk/ZkClientTest.scala @@ -1,26 +1,23 @@ package com.twitter.zk -import scala.collection.JavaConverters._ import scala.collection.Set +import scala.jdk.CollectionConverters._ import org.apache.zookeeper._ import org.apache.zookeeper.data.{ACL, Stat} -import org.junit.runner.RunWith -import org.mockito.Matchers.{eq => meq, _} +import org.mockito.ArgumentMatchers.{eq => meq, _} import org.mockito.Mockito._ import org.mockito._ import org.mockito.invocation.InvocationOnMock import org.mockito.stubbing.Answer -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.logging.{Level, Logger} import com.twitter.util._ +import org.scalatest.wordspec.AnyWordSpec -@RunWith(classOf[JUnitRunner]) -class ZkClientTest extends WordSpec with MockitoSugar { +class ZkClientTest extends AnyWordSpec with MockitoSugar { Logger.get("").setLevel(Level.FATAL) implicit val javaTimer = new JavaTimer(true) @@ -58,7 +55,9 @@ class ZkClientTest extends WordSpec with MockitoSugar { data: Array[Byte] = "".getBytes, acls: Seq[ACL] = zkClient.acl, mode: CreateMode = zkClient.mode - )(wait: => Future[String]): Unit = { + )( + wait: => Future[String] + ): Unit = { when( zk.create( meq(path), @@ -79,7 +78,12 @@ class ZkClientTest extends WordSpec with MockitoSugar { } def delete(path: String, version: Int)(wait: => Future[Unit]): Unit = { - when(zk.delete(meq(path), meq(version), any[AsyncCallback.VoidCallback], meq(null))) thenAnswer answer[ + when( + zk.delete( + meq(path), + meq(version), + any[AsyncCallback.VoidCallback], + meq(null))) thenAnswer answer[ AsyncCallback.VoidCallback ](2) { cbValue => wait onSuccess { _ => @@ -129,7 +133,12 @@ class ZkClientTest extends WordSpec with MockitoSugar { def getChildren(path: String)(children: => Future[ZNode.Children]): Unit = { val cb = ArgumentCaptor.forClass(classOf[AsyncCallback.Children2Callback]) - when(zk.getChildren(meq(path), meq(false), any[AsyncCallback.Children2Callback], meq(null))) thenAnswer answer[ + when( + zk.getChildren( + meq(path), + meq(false), + any[AsyncCallback.Children2Callback], + meq(null))) thenAnswer answer[ AsyncCallback.Children2Callback ](2) { cbValue => children onSuccess { znode => @@ -149,7 +158,11 @@ class ZkClientTest extends WordSpec with MockitoSugar { def watchChildren( path: String - )(children: => Future[ZNode.Children])(update: => Future[WatchedEvent]): Unit = { + )( + children: => Future[ZNode.Children] + )( + update: => Future[WatchedEvent] + ): Unit = { val w = ArgumentCaptor.forClass(classOf[Watcher]) val cb = ArgumentCaptor.forClass(classOf[AsyncCallback.Children2Callback]) when(zk.getChildren(meq(path), w.capture(), cb.capture(), meq(null))) thenAnswer answer[ @@ -169,7 +182,12 @@ class ZkClientTest extends WordSpec with MockitoSugar { } def getData(path: String)(result: => Future[ZNode.Data]): Unit = { - when(zk.getData(meq(path), meq(false), any[AsyncCallback.DataCallback], meq(null))) thenAnswer answer[ + when( + zk.getData( + meq(path), + meq(false), + any[AsyncCallback.DataCallback], + meq(null))) thenAnswer answer[ AsyncCallback.DataCallback ](2) { cbValue => result onSuccess { z => @@ -181,7 +199,13 @@ class ZkClientTest extends WordSpec with MockitoSugar { } } - def watchData(path: String)(result: => Future[ZNode.Data])(update: => Future[WatchedEvent]): Unit = { + def watchData( + path: String + )( + result: => Future[ZNode.Data] + )( + update: => Future[WatchedEvent] + ): Unit = { val w = ArgumentCaptor.forClass(classOf[Watcher]) val cb = ArgumentCaptor.forClass(classOf[AsyncCallback.DataCallback]) when(zk.getData(meq(path), w.capture(), cb.capture(), meq(null))) thenAnswer answer[ @@ -254,9 +278,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { i += 1 Future.exception(connectionLoss) } - .onSuccess { _ => - fail("Unexpected success") - } + .onSuccess { _ => fail("Unexpected success") } .handle { case e: KeeperException.ConnectionLossException => assert(e == connectionLoss) @@ -274,9 +296,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { i += 1 Future.Done } - .onSuccess { _ => - assert(i == 1) - } + .onSuccess { _ => assert(i == 1) } ) } @@ -290,9 +310,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { i += 1 throw rex } - .onSuccess { _ => - fail("Unexpected success") - } + .onSuccess { _ => fail("Unexpected success") } .handle { case e: RuntimeException => assert(e == rex) @@ -309,9 +327,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { i += 1 Future.exception(connectionLoss) } - .onSuccess { _ => - fail("Shouldn't have succeeded") - } + .onSuccess { _ => fail("Shouldn't have succeeded") } .handle { case e: KeeperException.ConnectionLossException => assert(i == 1) @@ -321,7 +337,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { } "transform" should { - val acl = ZooDefs.Ids.OPEN_ACL_UNSAFE.asScala + val acl = ZooDefs.Ids.OPEN_ACL_UNSAFE.asScala.toSeq val mode = CreateMode.EPHEMERAL_SEQUENTIAL val retries = 37 val retryPolicy = new RetryPolicy.Basic(retries) @@ -540,16 +556,10 @@ class ZkClientTest extends WordSpec with MockitoSugar { class MonitorHelper extends ExistHelper { val deleted = NodeEvent.Deleted(znode.path) def expectZNodes(n: Int): Unit = { - val results = 0 until n map { _ => - ZNode.Exists(znode, new Stat) - } - results foreach { r => - watch(znode.path)(Future(r.stat))(Future(deleted)) - } + val results = 0 until n map { _ => ZNode.Exists(znode, new Stat) } + results foreach { r => watch(znode.path)(Future(r.stat))(Future(deleted)) } val update = znode.exists.monitor() - results foreach { result => - assert(update.syncWait().get() == result) - } + results foreach { result => assert(update.syncWait().get() == result) } } } @@ -590,7 +600,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { ) offer.sync() intercept[KeeperException.NoNodeException] { - offer syncWait () get () + offer.syncWait().get } assert(offer.sync().isDefined == false) } @@ -675,7 +685,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { val update = znode.getChildren.monitor() results foreach { result => - val r = update syncWait () get () + val r = update.syncWait().get assert(r.path == result.path) } } @@ -744,9 +754,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { "In every generation there is a chosen one.", "She alone will stand against the vampires the demons and the forces of darkness.", "She is the slayer." - ) map { text => - ZNode.Data(znode, new Stat, text.getBytes) - } + ) map { text => ZNode.Data(znode, new Stat, text.getBytes) } results foreach { result => watchData(znode.path)(Future(result)) { Future(NodeEvent.ChildrenChanged(znode.path)) @@ -777,14 +785,8 @@ class ZkClientTest extends WordSpec with MockitoSugar { val treeRoot = ZNode.Children(zkClient("/arboreal"), new Stat, 'a' to 'e' map { _.toString }) val treeChildren = treeRoot +: ('a' to 'e') - .map { c => - ZNode.Children(treeRoot(c.toString), new Stat, 'a' to c map { _.toString }) - } - .flatMap { z => - z +: z.children.map { c => - ZNode.Children(c, new Stat, Nil) - } - } + .map { c => ZNode.Children(treeRoot(c.toString), new Stat, 'a' to c map { _.toString }) } + .flatMap { z => z +: z.children.map { c => ZNode.Children(c, new Stat, Nil) } } // Lay out node updates for the tree: Add a 'z' node to all nodes named 'a' val updateTree = treeChildren.collect { @@ -809,13 +811,11 @@ class ZkClientTest extends WordSpec with MockitoSugar { // Create promises for each node in the tree -- satisfying a promise will fire a // ChildrenChanged event for its associated node. val updatePromises = treeChildren.map { _.path -> new Promise[WatchedEvent] }.toMap - treeChildren foreach { z => - watchChildren(z.path)(Future(z))(updatePromises(z.path)) - } + treeChildren foreach { z => watchChildren(z.path)(Future(z))(updatePromises(z.path)) } val offer = treeRoot.monitorTree() treeChildren foreach { _ => - val ztu = Await.result(offer sync (), 1.second) + val ztu = Await.result(offer.sync(), 1.second) val e = expectedByPath(ztu.parent.path) assert(ztu.parent.path == e.parent.path) assert(ztu.added.map { _.path } == e.added.map { _.path }.toSet) @@ -828,7 +828,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { updateTree foreach { z => updatePromises get (z.path) foreach { _.setValue(event(z.path)) } - val ztu = Await.result(offer sync (), 1.second) + val ztu = Await.result(offer.sync(), 1.second) val e = updatesByPath(z.path) assert(ztu.parent.path == e.parent.path) assert(ztu.added.map { _.path } == e.added.map { _.path }.toSet) @@ -840,9 +840,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { } "ok" in okUpdates { NodeEvent.ChildrenChanged(_) } - "be resilient to disconnect" in okUpdates { _ => - StateEvent.Disconnected() - } + "be resilient to disconnect" in okUpdates { _ => StateEvent.Disconnected() } "stop on session expiration" in { treeChildren foreach { z => @@ -855,7 +853,7 @@ class ZkClientTest extends WordSpec with MockitoSugar { val offer = treeRoot.monitorTree() intercept[TimeoutException] { treeChildren foreach { _ => - val ztu = Await.result(offer sync (), 1.second) + val ztu = Await.result(offer.sync(), 1.second) val e = expectedByPath(ztu.parent.path) assert(ztu.parent == e.parent) assert(ztu.added.map { _.path } == e.added.map { _.path }.toSet) @@ -904,12 +902,14 @@ class ZkClientTest extends WordSpec with MockitoSugar { sync(znode.path) { Future.exception(new KeeperException.SystemErrorException) } - Await.ready(znode.sync() map { _ => - fail("Unexpected success") - } handle { - case e: KeeperException.SystemErrorException => - assert(e.getPath == znode.path) - }, 1.second) + Await.ready( + znode.sync() map { _ => + fail("Unexpected success") + } handle { + case e: KeeperException.SystemErrorException => + assert(e.getPath == znode.path) + }, + 1.second) } } } diff --git a/util-zk/src/test/scala/com/twitter/zk/coordination/BUILD b/util-zk/src/test/scala/com/twitter/zk/coordination/BUILD new file mode 100644 index 0000000000..828d368269 --- /dev/null +++ b/util-zk/src/test/scala/com/twitter/zk/coordination/BUILD @@ -0,0 +1,22 @@ +junit_tests( + sources = ["*.scala"], + compiler_option_sets = ["fatal_warnings"], + platform = "java8", + strict_deps = True, + tags = [ + "bazel-compatible", + "non-exclusive", + ], + dependencies = [ + "3rdparty/jvm/junit", + "3rdparty/jvm/org/apache/zookeeper:zookeeper-client", + "3rdparty/jvm/org/mockito:mockito-core", + "3rdparty/jvm/org/scalatest", + "3rdparty/jvm/org/scalatestplus:junit", + "3rdparty/jvm/org/scalatestplus:mockito", + "util/util-core:util-core-util", + "util/util-core/src/main/scala/com/twitter/conversions", + "util/util-zk/src/main/scala/com/twitter/zk", + "util/util-zk/src/main/scala/com/twitter/zk/coordination", + ], +) diff --git a/util-zk/src/test/scala/com/twitter/zk/coordination/ShardCoordinatorTest.scala b/util-zk/src/test/scala/com/twitter/zk/coordination/ShardCoordinatorTest.scala index 8e69781ee0..d275ccd629 100644 --- a/util-zk/src/test/scala/com/twitter/zk/coordination/ShardCoordinatorTest.scala +++ b/util-zk/src/test/scala/com/twitter/zk/coordination/ShardCoordinatorTest.scala @@ -1,19 +1,16 @@ package com.twitter.zk.coordination -import scala.collection.JavaConverters._ +import scala.jdk.CollectionConverters._ import org.apache.zookeeper.ZooDefs.Ids.OPEN_ACL_UNSAFE -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar +import org.scalatestplus.mockito.MockitoSugar -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util.{Await, Future, JavaTimer} import com.twitter.zk.{NativeConnector, RetryPolicy, ZkClient} +import org.scalatest.wordspec.AnyWordSpec -@RunWith(classOf[JUnitRunner]) -class ShardCoordinatorTest extends WordSpec with MockitoSugar { +class ShardCoordinatorTest extends AnyWordSpec with MockitoSugar { "ShardCoordinator" should { @@ -25,7 +22,7 @@ class ShardCoordinatorTest extends WordSpec with MockitoSugar { val connector = NativeConnector(connectString, 5.seconds, 10.minutes) val zk = ZkClient(connector) .withRetryPolicy(RetryPolicy.Basic(3)) - .withAcl(OPEN_ACL_UNSAFE.asScala) + .withAcl(OPEN_ACL_UNSAFE.asScala.toSeq) Await.result(Future { f(zk) } ensure { zk.release }) } diff --git a/util-zk/src/test/scala/com/twitter/zk/coordination/ZkAsyncSemaphoreTest.scala b/util-zk/src/test/scala/com/twitter/zk/coordination/ZkAsyncSemaphoreTest.scala index f80d2581eb..806545e6c1 100644 --- a/util-zk/src/test/scala/com/twitter/zk/coordination/ZkAsyncSemaphoreTest.scala +++ b/util-zk/src/test/scala/com/twitter/zk/coordination/ZkAsyncSemaphoreTest.scala @@ -1,23 +1,20 @@ package com.twitter.zk.coordination import com.twitter.concurrent.Permit -import com.twitter.conversions.time._ +import com.twitter.conversions.DurationOps._ import com.twitter.util._ import com.twitter.zk.{NativeConnector, RetryPolicy, ZkClient, ZNode} import java.util.concurrent.ConcurrentLinkedQueue import org.apache.zookeeper.KeeperException.NoNodeException import org.apache.zookeeper.ZooDefs.Ids.OPEN_ACL_UNSAFE -import org.junit.runner.RunWith -import org.scalatest.WordSpec -import org.scalatest.concurrent.AsyncAssertions +import org.scalatest.concurrent.Waiters import org.scalatest.concurrent.Eventually._ -import org.scalatest.junit.JUnitRunner -import org.scalatest.mockito.MockitoSugar import org.scalatest.time.{Millis, Seconds, Span} -import scala.collection.JavaConverters._ +import org.scalatestplus.mockito.MockitoSugar +import scala.jdk.CollectionConverters._ +import org.scalatest.wordspec.AnyWordSpec -@RunWith(classOf[JUnitRunner]) -class ZkAsyncSemaphoreTest extends WordSpec with MockitoSugar with AsyncAssertions { +class ZkAsyncSemaphoreTest extends AnyWordSpec with MockitoSugar with Waiters { "ZkAsyncSemaphore" should { @@ -30,15 +27,13 @@ class ZkAsyncSemaphoreTest extends WordSpec with MockitoSugar with AsyncAssertio val connector = NativeConnector(connectString, 5.seconds, 10.minutes) val zk = ZkClient(connector) .withRetryPolicy(RetryPolicy.Basic(3)) - .withAcl(OPEN_ACL_UNSAFE.asScala) + .withAcl(OPEN_ACL_UNSAFE.asScala.toSeq) Await.result(Future { f(zk) } ensure { zk.release }) } def acquire(sem: ZkAsyncSemaphore) = { - sem.acquire() onSuccess { permit => - permits add permit - } + sem.acquire() onSuccess { permit => permits add permit } } "provide a shared 2-permit semaphore and" should {