Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/daily_spend_bulk_upsert.py: 87%

67 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""One multi-row ``INSERT ... ON CONFLICT DO UPDATE`` per batch of daily spend rows. 

2 

3Emitting a statement per aggregated key put every replica's flush on the database as 

4hundreds of separate statements against the same handful of hot rows, each holding its 

5row locks for the rest of the enclosing batch transaction. Folding a batch into a single 

6statement keeps the aggregation identical while collapsing both the statement count and 

7the window in which those locks are held. 

8""" 

9 

10import uuid 

11from collections.abc import Mapping, Sequence 

12from dataclasses import dataclass 

13from itertools import groupby 

14from types import MappingProxyType 

15from typing import Final, Literal 

16 

17from pydantic import TypeAdapter 

18 

19DailySpendEntity = Literal["user", "team", "org", "tag", "end_user", "agent"] 

20 

21SqlValue = str | int | float | None 

22 

23# A queued daily spend transaction, read by column name because the columns are data 

24# here rather than literals. The concrete TypedDicts in _types.py all satisfy this. 

25SpendRow = Mapping[str, object] 

26 

27 

28@dataclass(frozen=True, slots=True) 

29class DailySpendTable: 

30 """The physical table behind one entity's daily rollup.""" 

31 

32 name: str 

33 entity_id_column: str 

34 carries_request_id: bool = False 

35 

36 

37DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType( 

38 { 

39 "user": DailySpendTable(name="LiteLLM_DailyUserSpend", entity_id_column="user_id"), 

40 "team": DailySpendTable(name="LiteLLM_DailyTeamSpend", entity_id_column="team_id"), 

41 "org": DailySpendTable(name="LiteLLM_DailyOrganizationSpend", entity_id_column="organization_id"), 

42 "end_user": DailySpendTable(name="LiteLLM_DailyEndUserSpend", entity_id_column="end_user_id"), 

43 "agent": DailySpendTable(name="LiteLLM_DailyAgentSpend", entity_id_column="agent_id"), 

44 "tag": DailySpendTable(name="LiteLLM_DailyTagSpend", entity_id_column="tag", carries_request_id=True), 

45 } 

46) 

47 

48_ENTITY_INPUT_KEYS: Final[Mapping[DailySpendEntity, str]] = MappingProxyType( 

49 { 

50 "user": "user", 

51 "team": "team_id", 

52 "org": "organization_id", 

53 "end_user": "end_user", 

54 "agent": "agent_id", 

55 "tag": "request_tags", 

56 } 

57) 

58_TAGS: Final = TypeAdapter(tuple[str, ...]) 

59 

60 

61def daily_spend_entity_ids(payload: Mapping[str, object], entity: DailySpendEntity) -> tuple[str | None, ...]: 

62 key: Final = _ENTITY_INPUT_KEYS[entity] 

63 if key not in payload: 63 ↛ 64line 63 didn't jump to line 64 because the condition on line 63 was never true

64 return () 

65 value: Final = payload[key] 

66 if entity == "tag": 

67 if value is None: 67 ↛ 68line 67 didn't jump to line 68 because the condition on line 67 was never true

68 return () 

69 tags: Final = _TAGS.validate_json(value) if isinstance(value, str) else _TAGS.validate_python(value) 

70 return tuple(dict.fromkeys(tags)) 

71 if value is None: 71 ↛ 72line 71 didn't jump to line 72 because the condition on line 71 was never true

72 return (None,) if entity == "user" else () 

73 if not isinstance(value, str) or (entity == "end_user" and not value): 73 ↛ 74line 73 didn't jump to line 74 because the condition on line 73 was never true

74 return () 

75 return (value,) 

76 

77 

78# The unique constraint's columns after the entity id, in constraint order. A NULL can 

79# never match itself in a unique index, so every one of these is normalized to '': the 

80# conflict target has to be NULL-free or the row is re-inserted on every single flush. 

81_KEY_COLUMNS: Final = ("date", "api_key", "model", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") 

82 

83_COUNTER_COLUMNS: Final = ( 

84 "prompt_tokens", 

85 "completion_tokens", 

86 "api_requests", 

87 "successful_requests", 

88 "failed_requests", 

89 "cache_read_input_tokens", 

90 "cache_creation_input_tokens", 

91 "compression_saved_tokens", 

92 "total_response_time_ms", 

93 "timed_requests", 

94) 

95_SPEND_COLUMNS: Final = ( 

96 "spend", 

97 "compression_savings_spend", 

98 "prompt_caching_savings_spend", 

99 "gateway_injected_caching_savings_spend", 

100 "autorouter_savings_spend", 

101) 

102 

103_CASTS: Final[Mapping[str, str]] = MappingProxyType( 

104 { 

105 **{column: "bigint" for column in _COUNTER_COLUMNS}, 

106 **{column: "double precision" for column in _SPEND_COLUMNS}, 

107 } 

108) 

109 

110 

111def _quoted(columns: Sequence[str]) -> str: 

112 return ", ".join(f'"{column}"' for column in columns) 

113 

114 

115def _as_text(value: object) -> str: 

116 return "" if value is None else str(value) 

117 

118 

119def _as_int(value: object) -> int: 

120 return int(value) if isinstance(value, (int, float)) else 0 

121 

122 

123def _as_float(value: object) -> float: 

124 return float(value) if isinstance(value, (int, float)) else 0.0 

125 

126 

127def conflict_key(table: DailySpendTable, transaction: SpendRow) -> tuple[str, ...]: 

128 """The tuple the database arbitrates the upsert on, normalized free of NULLs.""" 

129 return tuple(_as_text(transaction.get(column)) for column in (table.entity_id_column, *_KEY_COLUMNS)) 

130 

131 

132def _merge(group: Sequence[SpendRow]) -> SpendRow: 

133 if len(group) == 1: 133 ↛ 135line 133 didn't jump to line 135 because the condition on line 133 was always true

134 return group[0] 

135 return { 

136 **group[0], 

137 **{column: sum(_as_int(row.get(column)) for row in group) for column in _COUNTER_COLUMNS}, 

138 **{column: sum(_as_float(row.get(column)) for row in group) for column in _SPEND_COLUMNS}, 

139 } 

140 

141 

142def merge_by_conflict_key( 

143 table: DailySpendTable, 

144 transactions: Sequence[SpendRow], 

145) -> tuple[tuple[tuple[str, ...], SpendRow], ...]: 

146 """Batch entries keyed by the conflict tuple, in a deterministic order. 

147 

148 The queue keys transactions by their raw field values, so two entries differing only 

149 in a NULL versus an empty member reach the writer separately while arbitrating to the 

150 same row. Postgres rejects a statement whose values touch one row twice, so they are 

151 summed here into the single row they were always destined to become. Ordering by the 

152 key keeps concurrent writers taking row locks in the same sequence. 

153 """ 

154 ordered: Final = sorted(transactions, key=lambda transaction: conflict_key(table, transaction)) 

155 return tuple((key, _merge(tuple(group))) for key, group in groupby(ordered, key=lambda t: conflict_key(table, t))) 

156 

157 

158def _row_params( 

159 table: DailySpendTable, 

160 key: tuple[str, ...], 

161 transaction: SpendRow, 

162) -> tuple[SqlValue, ...]: 

163 request_id: Final = transaction.get("request_id") 

164 return ( 

165 str(uuid.uuid4()), 

166 *key, 

167 None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")), 

168 *(_as_int(transaction.get(column)) for column in _COUNTER_COLUMNS), 

169 *(_as_float(transaction.get(column)) for column in _SPEND_COLUMNS), 

170 *((None if request_id is None else _as_text(request_id),) if table.carries_request_id else ()), 

171 ) 

172 

173 

174def _insert_columns(table: DailySpendTable) -> tuple[str, ...]: 

175 return ( 

176 "id", 

177 table.entity_id_column, 

178 *_KEY_COLUMNS, 

179 "model_group", 

180 *_COUNTER_COLUMNS, 

181 *_SPEND_COLUMNS, 

182 *(("request_id",) if table.carries_request_id else ()), 

183 ) 

184 

185 

186def build_bulk_upsert( 

187 table: DailySpendTable, 

188 batch: Sequence[tuple[tuple[str, ...], SpendRow]], 

189) -> tuple[str, tuple[SqlValue, ...]]: 

190 """The single statement writing one merged batch, plus its positional arguments.""" 

191 columns: Final = _insert_columns(table) 

192 quoted_table: Final = f'"{table.name}"' 

193 rows: Final = ", ".join( 

194 "(" 

195 + ", ".join( 

196 f"${row_index * len(columns) + offset + 1}::{_CASTS.get(column, 'text')}" 

197 for offset, column in enumerate(columns) 

198 ) 

199 + ", (NOW() AT TIME ZONE 'UTC'))" 

200 for row_index in range(len(batch)) 

201 ) 

202 increments: Final = ", ".join( 

203 f'"{column}" = {quoted_table}."{column}" + EXCLUDED."{column}"' 

204 for column in (*_COUNTER_COLUMNS, *_SPEND_COLUMNS) 

205 ) 

206 # request_id names one arbitrary contributing request, so an entry carrying none must 

207 # not blank out the one already recorded. 

208 request_id_update: Final = ( 

209 f', "request_id" = COALESCE(EXCLUDED."request_id", {quoted_table}."request_id")' 

210 if table.carries_request_id 

211 else "" 

212 ) 

213 sql: Final = ( 

214 f'INSERT INTO {quoted_table} ({_quoted(columns)}, "updated_at")\n' 

215 f"VALUES {rows}\n" 

216 f"ON CONFLICT ({_quoted((table.entity_id_column, *_KEY_COLUMNS))}) DO UPDATE SET\n" 

217 f" {increments}{request_id_update},\n" 

218 f" \"updated_at\" = (NOW() AT TIME ZONE 'UTC')" 

219 ) 

220 return sql, tuple(value for key, transaction in batch for value in _row_params(table, key, transaction))