File size: 7,532 Bytes
020ed68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3698ef0
 
 
 
 
020ed68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
185dcd4
 
 
 
 
020ed68
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
import { Db, MongoClient } from 'mongodb'

/**
 * Spreads each data class across N Atlas clusters.
 *
 *   MONGO_URLS_LAMBDA  = "mongodb+srv://...,mongodb+srv://...,..."
 *   MONGO_URLS_COSTING = "mongodb+srv://...,..."
 *
 * Both fall back to MONGO_DB_URI, so a single-cluster setup keeps working.
 *
 * Writes use fill-and-spill: everything goes to the first shard with room, and
 * rolls to the next when that one nears quota. Hashing by func_name was the
 * obvious alternative and is wrong here - one function is ~99% of lambda traffic,
 * so it would pin nearly all writes to a single shard while the rest sat empty.
 *
 * Reads fan out to every shard and merge, because fill-and-spill means older
 * data lives on earlier shards.
 */

export type ShardGroup = 'lambda' | 'costing'

// Atlas M0 is 512MB. Roll over early so a burst can't hit the wall mid-batch.
const QUOTA_BYTES = Number(process.env.SHARD_QUOTA_BYTES ?? 512 * 1024 * 1024)
const ROLLOVER_AT = Number(process.env.SHARD_ROLLOVER_AT ?? 0.9)
const SIZE_TTL_MS = 5 * 60_000

/**
 * Every shard is opened with an EXPLICIT database name rather than whatever the
 * URI happens to carry. Connection strings that end in "/" (no db path) make the
 * driver fall back to "test", which would silently scatter data across a
 * different database on each cluster.
 */
const DB_NAME = process.env.MONGO_DB_NAME ?? 'costingDB'

/**
 * Env var names for a shard group, ordered by their number.
 *
 * Matched case-INSENSITIVELY on purpose: GitHub Actions uppercases every secret
 * name, so `CostingDBURL1` arrives as `COSTINGDBURL1`. A case-sensitive match
 * finds nothing there and the server starts with no clusters at all.
 *
 * Sorted numerically, not lexically, so 10 lands after 9 rather than after 1 -
 * fill-and-spill writes to the first shard with room, so order IS the policy.
 */
export function shardEnvKeys(prefix: string, env: Record<string, string | undefined>): string[] {
  const re = new RegExp(`^${prefix}(\\d+)$`, 'i')
  return Object.keys(env)
    .map((k) => ({ k, n: k.match(re)?.[1] }))
    .filter((x) => x.n !== undefined && env[x.k]?.trim())
    .sort((a, b) => Number(a.n) - Number(b.n))
    .map((x) => x.k)
}

function urisFor(group: ShardGroup): string[] {
  const prefix = group === 'lambda' ? 'LambdaDBURL' : 'CostingDBURL'

  const numbered = shardEnvKeys(prefix, process.env).map((k) => process.env[k]!.trim())
  if (numbered.length) return numbered

  // Fallbacks: one comma-separated var, then the original single-cluster setting.
  const raw =
    process.env[group === 'lambda' ? 'MONGO_URLS_LAMBDA' : 'MONGO_URLS_COSTING'] ??
    process.env.MONGO_DB_URI ??
    ''
  const uris = raw.split(',').map((u) => u.trim()).filter(Boolean)
  if (!uris.length) throw new Error(`no Mongo URLs configured for "${group}" (expected ${prefix}1, ${prefix}2, …)`)
  return uris
}

// Cached on globalThis so `next dev`'s module reloading doesn't leak a new pool
// per edit - same reason the original single-client helper did it.
declare global {
  // eslint-disable-next-line no-var
  var _mongoPool: Map<string, Promise<MongoClient>> | undefined
  // eslint-disable-next-line no-var
  var _mongoSizes: Map<ShardGroup, { at: number; used: number[] }> | undefined
}
const pool = (globalThis._mongoPool ??= new Map())
const sizes = (globalThis._mongoSizes ??= new Map())

function connect(uri: string): Promise<MongoClient> {
  let p = pool.get(uri)
  if (!p) {
    // Small pools, because the shard count is the multiplier here: at 50 clusters
    // the driver's default of 100 would allow 5,000 sockets from one process for
    // a workload that peaks in the low tens of ops/sec. Pools grow on demand, so
    // idle shards cost nothing.
    p = new MongoClient(uri, { maxPoolSize: 5 }).connect()
    pool.set(uri, p)
  }
  return p
}

export async function shards(group: ShardGroup): Promise<Db[]> {
  const clients = await Promise.all(urisFor(group).map(connect))
  return clients.map((c) => c.db(DB_NAME))
}

/** Bytes used per shard, cached - db.stats() on every write would be absurd. */
async function usage(group: ShardGroup): Promise<number[]> {
  const hit = sizes.get(group)
  if (hit && Date.now() - hit.at < SIZE_TTL_MS) return hit.used

  const dbs = await shards(group)
  const used = await Promise.all(
    dbs.map(async (db) => {
      try {
        const s = await db.command({ dbStats: 1, scale: 1 })
        // dataSize, NOT storageSize. Atlas bills uncompressed BSON plus indexes;
        // storageSize is the compressed on-disk figure and reads ~40% lower, so
        // using it reports a cluster as half full when Atlas has already blocked
        // writes on it.
        return (s.dataSize ?? 0) + (s.indexSize ?? 0)
      } catch {
        return 0 // unreachable shard: don't let a probe failure block writes
      }
    })
  )
  sizes.set(group, { at: Date.now(), used })
  return used
}

export function invalidateUsage(group: ShardGroup) {
  sizes.delete(group)
}

/** The shard writes should go to: first one under the rollover threshold. */
export async function writeShard(group: ShardGroup): Promise<Db> {
  const dbs = await shards(group)
  if (dbs.length === 1) return dbs[0]

  const used = await usage(group)
  const limit = QUOTA_BYTES * ROLLOVER_AT
  const i = used.findIndex((b) => b < limit)
  // Every shard full: use the last one and let Atlas raise the quota error
  // rather than silently dropping data.
  return dbs[i === -1 ? dbs.length - 1 : i]
}

export async function fanout<T>(group: ShardGroup, fn: (db: Db) => Promise<T>): Promise<T[]> {
  const dbs = await shards(group)
  return Promise.all(dbs.map(fn))
}

/** Sum a scalar that each shard reported independently. */
export function sumOf(results: number[]): number {
  return results.reduce((a, b) => a + b, 0)
}

/**
 * Merge per-shard $group output by _id, summing the named numeric fields.
 *
 * Callers MUST fan out without a $limit inside the pipeline and sort/limit via
 * this function instead: a top-10 taken per shard and then merged can miss an
 * entry that ranks 11th everywhere but first overall.
 */
export function mergeGroups<T extends { _id: unknown }>(
  perShard: T[][],
  fields: (keyof T)[],
  opts: { sortBy?: keyof T; limit?: number } = {}
): T[] {
  const merged = new Map<string, T>()

  for (const rows of perShard) {
    for (const row of rows) {
      const key = JSON.stringify(row._id)
      const hit = merged.get(key)
      if (!hit) {
        // Materialise every summed field so the output shape doesn't depend on
        // how many shards happened to be configured.
        const seed = { ...row }
        for (const f of fields) (seed[f] as number) = (seed[f] as number) ?? 0
        merged.set(key, seed)
        continue
      }
      for (const f of fields) {
        ;(hit[f] as number) = ((hit[f] as number) ?? 0) + ((row[f] as number) ?? 0)
      }
    }
  }

  let out = [...merged.values()]
  if (opts.sortBy) out.sort((a, b) => ((b[opts.sortBy!] as number) ?? 0) - ((a[opts.sortBy!] as number) ?? 0))
  if (opts.limit) out = out.slice(0, opts.limit)
  return out
}

/** Per-shard usage, for the admin status endpoint. */
export async function shardStatus(group: ShardGroup) {
  invalidateUsage(group)
  const used = await usage(group)
  return used.map((bytes, i) => ({
    shard: i,
    used_mb: +(bytes / 1e6).toFixed(1),
    quota_mb: +(QUOTA_BYTES / 1e6).toFixed(1),
    pct: +((bytes / QUOTA_BYTES) * 100).toFixed(1),
    accepting_writes: bytes < QUOTA_BYTES * ROLLOVER_AT,
  }))
}