-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel_manager.py
More file actions
203 lines (153 loc) · 7.14 KB
/
Copy pathmodel_manager.py
File metadata and controls
203 lines (153 loc) · 7.14 KB
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
201
202
import os
from datetime import datetime, timedelta
from zoneinfo import ZoneInfo
from collections import deque
import json
from logger import LogLevel, Logger
from pathlib import Path
logger = Logger() # Probably not great practice but its guarded with a mutex
DATA_FILE = "key_data.json"
class Limits:
def __init__(self, RPM, TPM, RPD):
self.RPM = RPM
self.TPM = TPM
self.RPD = RPD
# Update when changed, put preferred models higher
# https://aistudio.google.com/rate-limit?timeRange=last-28-days
model_limits = {
"gemini-3.5-flash": Limits(5, 250000, 20),
"gemini-3-flash-preview": Limits(5, 250000, 20),
"gemini-3.1-flash-lite": Limits(15, 250000, 500)
}
class Record:
def __init__(self, tokens_used, date):
self.tokens = tokens_used
self.date = date
self.RPM_error = False
self.TPM_error = False
self.RPD_error = False
self.DEMAND_error = False # High demand for the model
class ModelUsage:
def __init__(self, limits: Limits):
self.limits = limits
self.rolling_TPM = 0
self.RPD_Made = 0
# API limits reset at midnight
self.last_request_date = datetime.now(ZoneInfo("America/Los_Angeles")).date()
self.past_uses = deque()
def check_availability(self):
"""Call before reserving this model to see if you can"""
now = datetime.now(ZoneInfo("America/Los_Angeles"))
today = now.date()
while len(self.past_uses) > 0 and now - self.past_uses[0].date > timedelta(minutes=1, seconds=5):
record = self.past_uses.popleft()
self.rolling_TPM -= record.tokens
if self.last_request_date < today:
# We can reset our rate limits for today
self.RPD_Made = 0
self.last_request_date = today
if self.RPD_Made >= self.limits.RPD:
return False
if self.rolling_TPM >= self.limits.TPM:
return False
if len(self.past_uses) >= self.limits.RPM:
return False
return True
def reserve(self) -> Record:
"""Call the moment you commit to this key, before the request goes out."""
current_time = datetime.now(ZoneInfo("America/Los_Angeles"))
self.last_request_date = current_time.date()
record = Record(0, current_time)
self.past_uses.append(record)
self.RPD_Made += 1
return record
def finalize(self, record: Record, tokens_used):
"""Call once you know the real token count, after the response completes."""
if self.past_uses:
record.tokens = tokens_used
self.rolling_TPM += tokens_used
# We may have gotten an error also, our tracking isn't persistent between server saves
# and stuff can be finicky so if API returns an error we update to match it
if record.RPD_error:
self.RPD_Made = self.limits.RPD
if record.RPM_error:
for _ in range(self.limits.RPM):
self.past_uses.append(Record(0, datetime.now(ZoneInfo("America/Los_Angeles"))))
if record.TPM_error:
self.past_uses.append(Record(self.limits.TPM, datetime.now(ZoneInfo("America/Los_Angeles"))))
# Info for a specific key, has multiple models
class KeyInfo:
def __init__(self):
self.model_usages = {model: ModelUsage(limit) for model, limit in model_limits.items()}
def has_model_available(self, model):
return self.model_usages[model].check_availability()
def reserve_model(self, model) -> Record:
return self.model_usages[model].reserve()
def finalize(self, model, record: Record, tokens_used):
self.model_usages[model].finalize(record, tokens_used)
class APIRecord:
def __init__(self, api_key, model, record: Record):
self.key = api_key
self.model = model
self.record = record
class ModelManager:
# If time since now - last high demand timeout < back_off delta we don't use the model
BACKOFF_DELTA = timedelta(minutes=1)
def __init__(self):
self.key_infos = {key : KeyInfo() for id, key in os.environ.items() if "KEY_" in id}
self.load()
# We init with a sentinel value
self.last_model_overload = {model : datetime.fromisoformat("2000-01-01") for model in model_limits.keys()}
def reserve_model(self, model) -> APIRecord:
for key, info in self.key_infos.items():
if not info.has_model_available(model):
continue
return APIRecord(key, model, info.reserve_model(model))
# In this case no models are available of this type so throw
raise Exception("No keys have this model currently available")
""" Can return None if all models are exhausted """
def reserve_best_model(self) -> APIRecord:
for model in model_limits.keys():
logger.log(LogLevel.INFO, f"Trying keys for {model} model")
if datetime.now() - self.last_model_overload[model] < self.BACKOFF_DELTA:
logger.log(LogLevel.INFO, f"Backing off from {model} due to high demand")
continue
# The preferred models are inserted first
try:
return self.reserve_model(model)
except Exception:
# Give up if no keys are currently available for this model
# And try it with just a worse model
logger.log(LogLevel.INFO, f"Exhausted all keys for {model} model")
return None
def finalize(self, API_record: APIRecord, tokens_used):
if API_record.record.DEMAND_error:
# If there's a demand error we should note this down here first
self.last_model_overload[API_record.model] = datetime.now()
self.key_infos[API_record.key].finalize(API_record.model, API_record.record, tokens_used)
def save(self):
data = {}
for key, key_info in self.key_infos.items():
uses = {}
for model, usage in key_info.model_usages.items():
uses[model] = {
"RPD": usage.RPD_Made,
"last_request_date": usage.last_request_date.isoformat()
}
data[key] = uses
with open(DATA_FILE, "w", encoding="utf-8") as file:
json.dump(data, file, indent=4)
""" Call this after loading from env and it will overwrite some data with prior saved """
def load(self):
if not Path(DATA_FILE).is_file():
return
with open(DATA_FILE, "r") as file:
data = json.load(file)
for key, key_infos in data.items():
if key not in self.key_infos.keys(): continue
relevant_key_infos = self.key_infos[key]
for model, saved_info in key_infos.items():
if model not in relevant_key_infos.model_usages.keys(): continue
relevant_key_infos.model_usages[model].RPD_Made = int(saved_info["RPD"])
request_date = datetime.fromisoformat(saved_info["last_request_date"]).date()
relevant_key_infos.model_usages[model].last_request_date = request_date