1//===- AllocToken.cpp - Allocation Token Calculation ----------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// Definition of AllocToken modes and shared calculation of stateless token IDs.
10//
11//===----------------------------------------------------------------------===//
12
13#include "llvm/Support/AllocToken.h"
14#include "llvm/ADT/StringSwitch.h"
15#include "llvm/Support/ErrorHandling.h"
16#include "llvm/Support/MathExtras.h"
17#include "llvm/Support/SipHash.h"
18
19using namespace llvm;
20
21std::optional<AllocTokenMode>
22llvm::getAllocTokenModeFromString(StringRef Name) {
23 return StringSwitch<std::optional<AllocTokenMode>>(Name)
24 .Case(S: "increment", Value: AllocTokenMode::Increment)
25 .Case(S: "random", Value: AllocTokenMode::Random)
26 .Case(S: "typehash", Value: AllocTokenMode::TypeHash)
27 .Case(S: "typehashpointersplit", Value: AllocTokenMode::TypeHashPointerSplit)
28 .Case(S: "typefunchash", Value: AllocTokenMode::TypeFuncHash)
29 .Case(S: "typefunchashpointersplit",
30 Value: AllocTokenMode::TypeFuncHashPointerSplit)
31 .Case(S: "default", Value: DefaultAllocTokenMode)
32 .Default(Value: std::nullopt);
33}
34
35StringRef llvm::getAllocTokenModeAsString(AllocTokenMode Mode) {
36 switch (Mode) {
37 case AllocTokenMode::Increment:
38 return "increment";
39 case AllocTokenMode::Random:
40 return "random";
41 case AllocTokenMode::TypeHash:
42 return "typehash";
43 case AllocTokenMode::TypeHashPointerSplit:
44 return "typehashpointersplit";
45 case AllocTokenMode::TypeFuncHash:
46 return "typefunchash";
47 case AllocTokenMode::TypeFuncHashPointerSplit:
48 return "typefunchashpointersplit";
49 }
50 llvm_unreachable("Unknown AllocTokenMode");
51}
52
53static uint64_t getStableHash(const AllocTokenMetadata &Metadata,
54 uint64_t MaxTokens) {
55 return getStableSipHash(Str: Metadata.TypeName) % MaxTokens;
56}
57
58/// Splits the Bits bits into: [pointer flag,] type name hash, function name
59/// hash. The function name hash gets Bits / 2 bits. Bits is k if MaxTokens is
60/// 2^k-1 (e.g. SIZE_MAX, tokens in [0, MaxTokens]), or Log2(MaxTokens)
61/// otherwise. With pointer split, the MSB is the pointer flag.
62static uint64_t getTypeFuncHash(const AllocTokenMetadata &Metadata,
63 uint64_t MaxTokens, bool PointerSplit) {
64 if (MaxTokens == 1)
65 return 0;
66 const unsigned Bits =
67 isMask_64(Value: MaxTokens) ? llvm::countr_one(Value: MaxTokens) : Log2_64(Value: MaxTokens);
68 const unsigned FuncBits = Bits / 2;
69 unsigned TypeBits = Bits - FuncBits;
70 uint64_t Token = 0;
71 if (PointerSplit) {
72 --TypeBits;
73 Token = uint64_t(Metadata.ContainsPointer) << (Bits - 1);
74 }
75 // An empty type name denotes an unknown type.
76 if (!Metadata.TypeName.empty())
77 Token |= (getStableSipHash(Str: Metadata.TypeName) &
78 maskTrailingOnes<uint64_t>(N: TypeBits))
79 << FuncBits;
80 Token |= getStableSipHash(Str: *Metadata.FunctionName) &
81 maskTrailingOnes<uint64_t>(N: FuncBits);
82 return Token;
83}
84
85std::optional<uint64_t> llvm::getAllocToken(AllocTokenMode Mode,
86 const AllocTokenMetadata &Metadata,
87 uint64_t MaxTokens) {
88 assert(MaxTokens && "Must provide non-zero max tokens");
89
90 switch (Mode) {
91 case AllocTokenMode::Increment:
92 case AllocTokenMode::Random:
93 // Stateful modes cannot be implemented as a pure function.
94 return std::nullopt;
95
96 case AllocTokenMode::TypeFuncHash:
97 case AllocTokenMode::TypeFuncHashPointerSplit:
98 if (!Metadata.FunctionName)
99 return std::nullopt;
100 return getTypeFuncHash(Metadata, MaxTokens,
101 PointerSplit: Mode == AllocTokenMode::TypeFuncHashPointerSplit);
102
103 case AllocTokenMode::TypeHash:
104 return getStableHash(Metadata, MaxTokens);
105
106 case AllocTokenMode::TypeHashPointerSplit: {
107 if (MaxTokens == 1)
108 return 0;
109 const uint64_t HalfTokens = MaxTokens / 2;
110 uint64_t Hash = getStableHash(Metadata, MaxTokens: HalfTokens);
111 if (Metadata.ContainsPointer)
112 Hash += HalfTokens;
113 return Hash;
114 }
115 }
116
117 llvm_unreachable("");
118}
119