pytensorforge 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (146) hide show
  1. cli.py +604 -0
  2. pytensorforge-0.1.0.dist-info/METADATA +103 -0
  3. pytensorforge-0.1.0.dist-info/RECORD +146 -0
  4. pytensorforge-0.1.0.dist-info/WHEEL +5 -0
  5. pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
  6. pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
  7. src/__init__.py +0 -0
  8. src/activations/Activation.py +4 -0
  9. src/activations/ELU.py +11 -0
  10. src/activations/GELU.py +6 -0
  11. src/activations/ReLU.py +27 -0
  12. src/activations/SELU.py +14 -0
  13. src/activations/Sigmoid.py +27 -0
  14. src/activations/Softmax.py +84 -0
  15. src/activations/Tanh.py +29 -0
  16. src/activations/__init__.py +17 -0
  17. src/config.py +120 -0
  18. src/core/Matrix.py +3 -0
  19. src/core/Scalar.py +18 -0
  20. src/core/Tensor.py +866 -0
  21. src/core/Vector.py +31 -0
  22. src/core/__init__.py +0 -0
  23. src/data/__init__.py +0 -0
  24. src/data/chat_dataset.py +188 -0
  25. src/data/corpus.py +104 -0
  26. src/data/document_stream.py +178 -0
  27. src/data/parallel_encode.py +86 -0
  28. src/data/prefetch.py +62 -0
  29. src/data/shard_builder.py +119 -0
  30. src/data/shard_writer.py +81 -0
  31. src/data/sharded_dataset.py +112 -0
  32. src/data/streaming_dataset.py +132 -0
  33. src/data/validation.py +212 -0
  34. src/inference/__init__.py +0 -0
  35. src/inference/chat_template.py +384 -0
  36. src/inference/config.py +48 -0
  37. src/inference/engine.py +241 -0
  38. src/inference/export.py +133 -0
  39. src/inference/kv_cache.py +65 -0
  40. src/inference/runtime.py +161 -0
  41. src/inference/sampling.py +42 -0
  42. src/inference/scheduler.py +473 -0
  43. src/inference/text.py +67 -0
  44. src/initializers/Constant.py +9 -0
  45. src/initializers/GlorotNormal.py +15 -0
  46. src/initializers/GlorotUniform.py +26 -0
  47. src/initializers/HeNormal.py +15 -0
  48. src/initializers/HeUniform.py +14 -0
  49. src/initializers/Initializer.py +4 -0
  50. src/initializers/LecunNormal.py +16 -0
  51. src/initializers/LecunUniform.py +14 -0
  52. src/initializers/Ones.py +6 -0
  53. src/initializers/Orthogonal.py +14 -0
  54. src/initializers/RandomNormal.py +14 -0
  55. src/initializers/RandomUniform.py +14 -0
  56. src/initializers/Zeros.py +8 -0
  57. src/initializers/__init__.py +17 -0
  58. src/loss/CategoricalCrossEntropy.py +9 -0
  59. src/loss/CrossEntropyLoss.py +34 -0
  60. src/loss/CrossEntropyWithLogitsLoss.py +59 -0
  61. src/loss/Hinge.py +5 -0
  62. src/loss/Huber.py +22 -0
  63. src/loss/Loss.py +6 -0
  64. src/loss/MSE.py +7 -0
  65. src/loss/MSELoss.py +10 -0
  66. src/loss/SparseCategoricalCrossEntropy.py +15 -0
  67. src/loss/__init__.py +18 -0
  68. src/loss/bce.py +34 -0
  69. src/loss/mae.py +16 -0
  70. src/math/__init__.py +0 -0
  71. src/math/clip.py +37 -0
  72. src/math/exp.py +27 -0
  73. src/math/log.py +25 -0
  74. src/math/sigmoid.py +5 -0
  75. src/models/__init__.py +0 -0
  76. src/models/embedding/Embedding.py +65 -0
  77. src/models/embedding/__init__.py +0 -0
  78. src/models/gpt/__init__.py +0 -0
  79. src/models/gpt/attention.py +158 -0
  80. src/models/gpt/block.py +74 -0
  81. src/models/gpt/config.py +103 -0
  82. src/models/gpt/context.py +44 -0
  83. src/models/gpt/model.py +165 -0
  84. src/models/gpt/recompute.py +35 -0
  85. src/models/gpt/rope.py +84 -0
  86. src/models/regression/Linear.py +51 -0
  87. src/models/regression/Logistic.py +36 -0
  88. src/models/regression/__init__.py +0 -0
  89. src/models/seq/Sequential.py +297 -0
  90. src/models/seq/__init__.py +0 -0
  91. src/models/svm/__init__.py +0 -0
  92. src/models/tokenizer/BPETokenizer.py +228 -0
  93. src/models/tokenizer/__init__.py +0 -0
  94. src/models/transformers/Dropout.py +35 -0
  95. src/models/transformers/LastToken.py +10 -0
  96. src/models/transformers/LayerNorm.py +54 -0
  97. src/models/transformers/Linear.py +18 -0
  98. src/models/transformers/MultiHeadAttention.py +130 -0
  99. src/models/transformers/TransformerBlock.py +79 -0
  100. src/models/transformers/__init__.py +0 -0
  101. src/neural/Dense.py +58 -0
  102. src/neural/LSTM.py +167 -0
  103. src/neural/Layer.py +72 -0
  104. src/neural/Parameter.py +30 -0
  105. src/neural/RNN.py +83 -0
  106. src/neural/__init__.py +0 -0
  107. src/ops/__init__.py +0 -0
  108. src/ops/stack.py +40 -0
  109. src/optimizers/Adagrad.py +31 -0
  110. src/optimizers/Adam.py +98 -0
  111. src/optimizers/AdamW.py +84 -0
  112. src/optimizers/Batch.py +11 -0
  113. src/optimizers/Nesterov.py +35 -0
  114. src/optimizers/Optimizer.py +18 -0
  115. src/optimizers/RMSProp.py +35 -0
  116. src/optimizers/SGD.py +30 -0
  117. src/optimizers/SGDMomentum.py +28 -0
  118. src/optimizers/__init__.py +9 -0
  119. src/scaling/StandardScaler.py +15 -0
  120. src/scaling/__init__.py +0 -0
  121. src/serialization/__init__.py +0 -0
  122. src/serialization/checkpoint.py +58 -0
  123. src/serialization/modelio.py +132 -0
  124. src/serving/__init__.py +0 -0
  125. src/serving/app.py +792 -0
  126. src/serving/config.py +216 -0
  127. src/serving/errors.py +51 -0
  128. src/serving/http.py +599 -0
  129. src/serving/metrics.py +293 -0
  130. src/serving/model_server.py +287 -0
  131. src/serving/protocol.py +377 -0
  132. src/serving/security.py +200 -0
  133. src/serving/server.py +121 -0
  134. src/tokenization/__init__.py +0 -0
  135. src/tokenization/base.py +75 -0
  136. src/tokenization/bpe.py +190 -0
  137. src/tokenization/bytebpe.py +476 -0
  138. src/tokenization/registry.py +28 -0
  139. src/training/__init__.py +0 -0
  140. src/training/checkpoint_manager.py +101 -0
  141. src/training/experiment.py +71 -0
  142. src/training/losses.py +42 -0
  143. src/training/precision.py +141 -0
  144. src/training/profiler.py +38 -0
  145. src/training/scheduler.py +50 -0
  146. src/training/trainer.py +594 -0
@@ -0,0 +1,377 @@
1
+ import time
2
+ import uuid
3
+ from dataclasses import dataclass, field
4
+ from typing import List, Optional, Union
5
+
6
+ from src.serving.errors import bad_request
7
+
8
+ CHAT_ROLES = ("system", "user", "assistant", "developer")
9
+
10
+ NEUTRAL_PARAMS = {
11
+ "n": 1,
12
+ "presence_penalty": 0,
13
+ "frequency_penalty": 0,
14
+ "logprobs": False,
15
+ "logit_bias": None,
16
+ "echo": False,
17
+ "best_of": 1,
18
+ "suffix": None,
19
+ "parallel_tool_calls": None,
20
+ "top_logprobs": None,
21
+ }
22
+
23
+ IGNORED_PARAMS = ("user", "metadata", "store", "service_tier", "response_format")
24
+
25
+ UNSUPPORTED_PARAMS = ("tools", "tool_choice", "functions", "function_call", "audio", "modalities", "prediction")
26
+
27
+
28
+ @dataclass
29
+ class SamplingParams:
30
+ max_tokens: Optional[int] = None
31
+ temperature: Optional[float] = None
32
+ top_p: Optional[float] = None
33
+ top_k: Optional[int] = None
34
+ repetition_penalty: Optional[float] = None
35
+ seed: Optional[int] = None
36
+ stop: List[str] = field(default_factory=list)
37
+
38
+
39
+ @dataclass
40
+ class ChatRequest:
41
+ model: str
42
+ messages: List[dict]
43
+ sampling: SamplingParams
44
+ stream: bool = False
45
+ include_usage: bool = False
46
+
47
+
48
+ @dataclass
49
+ class CompletionRequest:
50
+ model: str
51
+ prompt: Union[str, List[int]]
52
+ sampling: SamplingParams
53
+ stream: bool = False
54
+ include_usage: bool = False
55
+
56
+
57
+ def _is_int(v):
58
+ return isinstance(v, int) and not isinstance(v, bool)
59
+
60
+
61
+ def _is_num(v):
62
+ return (isinstance(v, (int, float))) and not isinstance(v, bool)
63
+
64
+
65
+ def _number(body, name, lo=None, hi=None, lo_open=False, integer=False):
66
+ if name not in body or body[name] is None:
67
+ return None
68
+
69
+ v = body[name]
70
+
71
+ if integer and not _is_int(v):
72
+ raise bad_request(f"'{name}' must be an integer", param=name)
73
+
74
+ if not integer and not _is_num(v):
75
+ raise bad_request(f"'{name}' must be a number", param=name)
76
+
77
+ if lo is not None and (v <= lo if lo_open else v < lo):
78
+ raise bad_request(f"'{name}' must be {'>' if lo_open else '>='} {lo}", param=name)
79
+
80
+ if hi is not None and v > hi:
81
+ raise bad_request(f"'{name}' must be <= {hi}", param=name)
82
+
83
+ return v
84
+
85
+
86
+ def _check_params(body, allowed):
87
+ if not isinstance(body, dict):
88
+ raise bad_request("request body must be a JSON object")
89
+
90
+ for name in UNSUPPORTED_PARAMS:
91
+ if body.get(name) not in (None, [], {}):
92
+ raise bad_request(f"'{name}' is not supported by this server", param=name, code="unsupported_parameter")
93
+
94
+ for name, neutral in NEUTRAL_PARAMS.items():
95
+ if name in body and body[name] is not None and body[name] != neutral:
96
+ raise bad_request(
97
+ f"'{name}'={body[name]!r} is not supported; only {neutral!r} is accepted",
98
+ param=name, code="unsupported_parameter",
99
+ )
100
+
101
+ known = set(allowed) | set(NEUTRAL_PARAMS) | set(IGNORED_PARAMS) | set(UNSUPPORTED_PARAMS)
102
+ unknown = sorted(k for k in body if k not in known)
103
+
104
+ if unknown:
105
+ raise bad_request(f"unrecognized request parameter(s): {unknown}", param=unknown[0])
106
+
107
+
108
+ def _model(body):
109
+ model = body.get("model")
110
+
111
+ if not isinstance(model, str) or not model:
112
+ raise bad_request("'model' is required and must be a string", param="model")
113
+
114
+ return model
115
+
116
+
117
+ def _stop(body, limits):
118
+ stop = body.get("stop")
119
+
120
+ if stop is None:
121
+ return []
122
+
123
+ if isinstance(stop, str):
124
+ stop = [stop]
125
+
126
+ if not isinstance(stop, list) or not all(isinstance(s, str) for s in stop):
127
+ raise bad_request("'stop' must be a string or a list of strings", param="stop")
128
+
129
+ stop = [s for s in stop if s]
130
+
131
+ if len(stop) > limits.max_stop_sequences:
132
+ raise bad_request(f"at most {limits.max_stop_sequences} stop sequences are allowed", param="stop")
133
+
134
+ if any(len(s) > limits.max_stop_sequence_chars for s in stop):
135
+ raise bad_request(f"stop sequences must be at most {limits.max_stop_sequence_chars} characters", param="stop")
136
+
137
+ return stop
138
+
139
+
140
+ def _stream(body):
141
+ stream = body.get("stream", False)
142
+
143
+ if stream is None:
144
+ stream = False
145
+
146
+ if not isinstance(stream, bool):
147
+ raise bad_request("'stream' must be a boolean", param="stream")
148
+
149
+ opts = body.get("stream_options")
150
+ include_usage = False
151
+
152
+ if opts is not None:
153
+ if not isinstance(opts, dict) or set(opts) - {"include_usage"}:
154
+ raise bad_request("'stream_options' supports only 'include_usage'", param="stream_options")
155
+
156
+ if not stream:
157
+ raise bad_request("'stream_options' is only allowed when 'stream' is true", param="stream_options")
158
+
159
+ include_usage = bool(opts.get("include_usage", False))
160
+
161
+ return stream, include_usage
162
+
163
+
164
+ def _sampling(body, limits, token_field_names):
165
+ max_tokens = None
166
+ source = None
167
+
168
+ for name in token_field_names:
169
+ v = _number(body, name, lo=1, integer=True)
170
+
171
+ if v is not None:
172
+ if max_tokens is not None and v != max_tokens:
173
+ raise bad_request(f"conflicting values for {' and '.join(token_field_names)}", param=name)
174
+ max_tokens = v
175
+ source = source or name
176
+
177
+ if max_tokens is not None and max_tokens > limits.max_generation_tokens:
178
+ raise bad_request(
179
+ f"{source}={max_tokens} exceeds this server's limit of {limits.max_generation_tokens}",
180
+ param=source, code="max_tokens_exceeded",
181
+ )
182
+
183
+ return SamplingParams(
184
+ max_tokens=max_tokens,
185
+ temperature=_number(body, "temperature", lo=0, hi=2),
186
+ top_p=_number(body, "top_p", lo=0, hi=1, lo_open=True),
187
+ top_k=_number(body, "top_k", lo=0, integer=True),
188
+ repetition_penalty=_number(body, "repetition_penalty", lo=0, hi=10, lo_open=True),
189
+ seed=_number(body, "seed", integer=True),
190
+ stop=_stop(body, limits),
191
+ )
192
+
193
+
194
+ def _content_text(content, idx):
195
+ if isinstance(content, str):
196
+ return content
197
+
198
+ if isinstance(content, list):
199
+ parts = []
200
+
201
+ for part in content:
202
+ if not isinstance(part, dict) or part.get("type") != "text" or not isinstance(part.get("text"), str):
203
+ raise bad_request(
204
+ f"messages[{idx}].content: only text content parts are supported", param=f"messages[{idx}].content"
205
+ )
206
+ parts.append(part["text"])
207
+
208
+ return "".join(parts)
209
+
210
+ if content is None:
211
+ return ""
212
+
213
+ raise bad_request(f"messages[{idx}].content must be a string or a list of text parts",
214
+ param=f"messages[{idx}].content")
215
+
216
+
217
+ CHAT_FIELDS = ("model", "messages", "stream", "stream_options", "max_tokens", "max_completion_tokens",
218
+ "temperature", "top_p", "top_k", "repetition_penalty", "seed", "stop")
219
+
220
+ COMPLETION_FIELDS = ("model", "prompt", "stream", "stream_options", "max_tokens", "temperature", "top_p",
221
+ "top_k", "repetition_penalty", "seed", "stop")
222
+
223
+
224
+ def parse_chat_request(body, limits):
225
+ _check_params(body, CHAT_FIELDS)
226
+
227
+ messages = body.get("messages")
228
+
229
+ if not isinstance(messages, list) or not messages:
230
+ raise bad_request("'messages' must be a non-empty list", param="messages")
231
+
232
+ if len(messages) > limits.max_messages:
233
+ raise bad_request(f"at most {limits.max_messages} messages are allowed", param="messages")
234
+
235
+ out = []
236
+ total_chars = 0
237
+
238
+ for i, m in enumerate(messages):
239
+ if not isinstance(m, dict):
240
+ raise bad_request(f"messages[{i}] must be an object", param=f"messages[{i}]")
241
+
242
+ role = m.get("role")
243
+
244
+ if role not in CHAT_ROLES:
245
+ raise bad_request(f"messages[{i}].role must be one of {list(CHAT_ROLES)}", param=f"messages[{i}].role")
246
+
247
+ if m.get("tool_calls") or role == "tool":
248
+ raise bad_request("tool messages are not supported", param=f"messages[{i}]")
249
+
250
+ text = _content_text(m.get("content"), i)
251
+ total_chars += len(text)
252
+ out.append({"role": "system" if role == "developer" else role, "content": text})
253
+
254
+ if total_chars > limits.max_prompt_chars:
255
+ raise bad_request(f"messages exceed the {limits.max_prompt_chars}-character limit",
256
+ param="messages", code="prompt_too_large")
257
+
258
+ if out[-1]["role"] == "assistant":
259
+ raise bad_request("the last message must not be from the assistant", param="messages")
260
+
261
+ stream, include_usage = _stream(body)
262
+
263
+ return ChatRequest(
264
+ model=_model(body),
265
+ messages=out,
266
+ sampling=_sampling(body, limits, ("max_tokens", "max_completion_tokens")),
267
+ stream=stream,
268
+ include_usage=include_usage,
269
+ )
270
+
271
+
272
+ def parse_completion_request(body, limits):
273
+ _check_params(body, COMPLETION_FIELDS)
274
+
275
+ prompt = body.get("prompt")
276
+
277
+ if isinstance(prompt, list) and len(prompt) == 1 and isinstance(prompt[0], str):
278
+ prompt = prompt[0]
279
+ elif isinstance(prompt, list) and len(prompt) == 1 and isinstance(prompt[0], list):
280
+ prompt = prompt[0]
281
+
282
+ if isinstance(prompt, str):
283
+ if len(prompt) > limits.max_prompt_chars:
284
+ raise bad_request(f"prompt exceeds the {limits.max_prompt_chars}-character limit",
285
+ param="prompt", code="prompt_too_large")
286
+ elif isinstance(prompt, list) and prompt and all(_is_int(t) for t in prompt):
287
+ if len(prompt) > limits.max_prompt_chars:
288
+ raise bad_request("prompt is too long", param="prompt", code="prompt_too_large")
289
+ elif isinstance(prompt, list):
290
+ raise bad_request("batched prompts are not supported; send one prompt per request", param="prompt")
291
+ else:
292
+ raise bad_request("'prompt' must be a string or a list of token ids", param="prompt")
293
+
294
+ stream, include_usage = _stream(body)
295
+
296
+ return CompletionRequest(
297
+ model=_model(body),
298
+ prompt=prompt,
299
+ sampling=_sampling(body, limits, ("max_tokens",)),
300
+ stream=stream,
301
+ include_usage=include_usage,
302
+ )
303
+
304
+
305
+ def new_id(prefix):
306
+ return f"{prefix}-{uuid.uuid4().hex[:24]}"
307
+
308
+
309
+ def usage_dict(prompt_tokens, completion_tokens):
310
+ return {
311
+ "prompt_tokens": prompt_tokens,
312
+ "completion_tokens": completion_tokens,
313
+ "total_tokens": prompt_tokens + completion_tokens,
314
+ }
315
+
316
+
317
+ def chat_response(rid, model, fingerprint, text, finish_reason, usage):
318
+ return {
319
+ "id": rid,
320
+ "object": "chat.completion",
321
+ "created": int(time.time()),
322
+ "model": model,
323
+ "system_fingerprint": fingerprint,
324
+ "choices": [{
325
+ "index": 0,
326
+ "message": {"role": "assistant", "content": text},
327
+ "logprobs": None,
328
+ "finish_reason": finish_reason,
329
+ }],
330
+ "usage": usage,
331
+ }
332
+
333
+
334
+ def chat_chunk(rid, created, model, fingerprint, delta, finish_reason=None, usage=None, include_choice=True):
335
+ chunk = {
336
+ "id": rid,
337
+ "object": "chat.completion.chunk",
338
+ "created": created,
339
+ "model": model,
340
+ "system_fingerprint": fingerprint,
341
+ "choices": [{"index": 0, "delta": delta, "logprobs": None, "finish_reason": finish_reason}]
342
+ if include_choice else [],
343
+ }
344
+
345
+ if usage is not None:
346
+ chunk["usage"] = usage
347
+
348
+ return chunk
349
+
350
+
351
+ def completion_response(rid, model, fingerprint, text, finish_reason, usage):
352
+ return {
353
+ "id": rid,
354
+ "object": "text_completion",
355
+ "created": int(time.time()),
356
+ "model": model,
357
+ "system_fingerprint": fingerprint,
358
+ "choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish_reason}],
359
+ "usage": usage,
360
+ }
361
+
362
+
363
+ def completion_chunk(rid, created, model, fingerprint, text, finish_reason=None, usage=None, include_choice=True):
364
+ chunk = {
365
+ "id": rid,
366
+ "object": "text_completion",
367
+ "created": created,
368
+ "model": model,
369
+ "system_fingerprint": fingerprint,
370
+ "choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish_reason}]
371
+ if include_choice else [],
372
+ }
373
+
374
+ if usage is not None:
375
+ chunk["usage"] = usage
376
+
377
+ return chunk
@@ -0,0 +1,200 @@
1
+ import hashlib
2
+ import hmac
3
+ import ipaddress
4
+ import os
5
+ import time
6
+ from collections import OrderedDict
7
+ from dataclasses import dataclass
8
+
9
+ from src.serving import errors
10
+
11
+
12
+ def hash_key(key):
13
+ return hashlib.sha256(key.encode("utf-8")).hexdigest()
14
+
15
+
16
+ def _parse_entry(entry):
17
+ entry = entry.strip()
18
+
19
+ if not entry or entry.startswith("#"):
20
+ return None
21
+
22
+ if entry.startswith("sha256:"):
23
+ digest = entry[len("sha256:"):].strip().lower()
24
+
25
+ if len(digest) != 64 or any(c not in "0123456789abcdef" for c in digest):
26
+ raise ValueError("malformed sha256 key entry")
27
+
28
+ return digest
29
+
30
+ if len(entry) < 16:
31
+ raise ValueError("API keys must be at least 16 characters")
32
+
33
+ return hash_key(entry)
34
+
35
+
36
+ def _read_key_file(path):
37
+ with open(path) as f:
38
+ return [line for line in f.read().splitlines()]
39
+
40
+
41
+ def _collect(inline, path, env_name):
42
+ raw = list(inline or [])
43
+
44
+ if path:
45
+ raw.extend(_read_key_file(path))
46
+
47
+ if env_name and os.environ.get(env_name):
48
+ raw.extend(os.environ[env_name].split(","))
49
+
50
+ return [d for d in (_parse_entry(e) for e in raw) if d]
51
+
52
+
53
+ @dataclass(frozen=True)
54
+ class Principal:
55
+ id: str
56
+ authenticated: bool
57
+ admin: bool = False
58
+
59
+
60
+ class KeyStore:
61
+
62
+ def __init__(self, user_digests, admin_digests, allow_unauthenticated=False):
63
+ self._user = [bytes.fromhex(d) for d in dict.fromkeys(user_digests)]
64
+ self._admin = [bytes.fromhex(d) for d in dict.fromkeys(admin_digests)]
65
+ self.allow_unauthenticated = allow_unauthenticated
66
+
67
+ @classmethod
68
+ def from_config(cls, sec):
69
+ users = _collect(sec.api_keys, sec.api_keys_file, sec.api_keys_env)
70
+ admins = _collect(sec.admin_keys, sec.admin_keys_file, None)
71
+ return cls(users, admins, sec.allow_unauthenticated)
72
+
73
+ @property
74
+ def enabled(self):
75
+ return bool(self._user or self._admin)
76
+
77
+ @staticmethod
78
+ def extract(headers):
79
+ auth = headers.get("authorization", "")
80
+
81
+ if auth[:7].lower() == "bearer ":
82
+ return auth[7:].strip() or None
83
+
84
+ key = headers.get("x-api-key", "").strip()
85
+
86
+ return key or None
87
+
88
+ @staticmethod
89
+ def _match(digest, pool):
90
+ found = False
91
+
92
+ for candidate in pool:
93
+ found |= hmac.compare_digest(digest, candidate)
94
+
95
+ return found
96
+
97
+ def authenticate(self, headers, client_ip):
98
+ key = self.extract(headers)
99
+
100
+ if key is None:
101
+ if self.enabled and not self.allow_unauthenticated:
102
+ raise errors.unauthorized("missing API key; send 'Authorization: Bearer <key>'")
103
+
104
+ return Principal(f"ip:{client_ip}", authenticated=False)
105
+
106
+ digest = hashlib.sha256(key.encode("utf-8")).digest()
107
+ is_admin = self._match(digest, self._admin)
108
+ is_user = self._match(digest, self._user)
109
+
110
+ if not (is_admin or is_user):
111
+ raise errors.unauthorized()
112
+
113
+ return Principal(f"key:{digest.hex()[:12]}", authenticated=True, admin=is_admin)
114
+
115
+
116
+ def is_loopback(host):
117
+ if host in ("localhost", ""):
118
+ return True
119
+
120
+ try:
121
+ return ipaddress.ip_address(host).is_loopback
122
+ except ValueError:
123
+ return False
124
+
125
+
126
+ class _Bucket:
127
+ __slots__ = ("tokens", "updated", "in_flight")
128
+
129
+ def __init__(self, capacity, now):
130
+ self.tokens = float(capacity)
131
+ self.updated = now
132
+ self.in_flight = 0
133
+
134
+
135
+ class RateLimiter:
136
+
137
+ def __init__(self, requests_per_minute, burst, max_concurrent, max_tracked=10_000, clock=time.monotonic):
138
+ self.rate = requests_per_minute / 60.0
139
+ self.capacity = burst
140
+ self.max_concurrent = max_concurrent
141
+ self.max_tracked = max_tracked
142
+ self.clock = clock
143
+ self._buckets = OrderedDict()
144
+
145
+ @classmethod
146
+ def from_config(cls, rl):
147
+ return cls(rl.requests_per_minute, rl.burst, rl.max_concurrent_per_principal, rl.max_tracked_principals)
148
+
149
+ def _bucket(self, key, now):
150
+ b = self._buckets.get(key)
151
+
152
+ if b is None:
153
+ b = _Bucket(self.capacity, now)
154
+ self._buckets[key] = b
155
+
156
+ while len(self._buckets) > self.max_tracked:
157
+ oldest, victim = next(iter(self._buckets.items()))
158
+
159
+ if victim.in_flight:
160
+ self._buckets.move_to_end(oldest)
161
+ break
162
+
163
+ self._buckets.popitem(last=False)
164
+ else:
165
+ self._buckets.move_to_end(key)
166
+
167
+ b.tokens = min(self.capacity, b.tokens + (now - b.updated) * self.rate)
168
+ b.updated = now
169
+
170
+ return b
171
+
172
+ def acquire(self, key):
173
+ now = self.clock()
174
+ b = self._bucket(key, now)
175
+
176
+ if b.in_flight >= self.max_concurrent:
177
+ raise errors.rate_limited(
178
+ f"too many concurrent requests for this API key (limit {self.max_concurrent})", 1
179
+ )
180
+
181
+ if b.tokens < 1.0:
182
+ raise errors.rate_limited("request rate limit exceeded", (1.0 - b.tokens) / self.rate)
183
+
184
+ b.tokens -= 1.0
185
+ b.in_flight += 1
186
+
187
+ def release(self, key):
188
+ b = self._buckets.get(key)
189
+
190
+ if b is not None and b.in_flight > 0:
191
+ b.in_flight -= 1
192
+
193
+ def charge(self, key):
194
+ now = self.clock()
195
+ b = self._bucket(key, now)
196
+
197
+ if b.tokens < 1.0:
198
+ raise errors.rate_limited("request rate limit exceeded", (1.0 - b.tokens) / self.rate)
199
+
200
+ b.tokens -= 1.0
src/serving/server.py ADDED
@@ -0,0 +1,121 @@
1
+ import asyncio
2
+ import logging
3
+ import os
4
+ import signal
5
+ import threading
6
+
7
+ from src.serving.app import APIApp
8
+ from src.serving.http import HTTPServer
9
+ from src.serving.model_server import ModelServer
10
+ from src.serving.security import is_loopback
11
+
12
+ log = logging.getLogger("ptf.server")
13
+
14
+ DEFAULT_UI_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))),
15
+ "web", "dist")
16
+
17
+
18
+ class InsecureConfiguration(ValueError):
19
+ pass
20
+
21
+
22
+ def check_security(config, keys):
23
+ if keys.enabled:
24
+ return
25
+
26
+ if config.security.allow_unauthenticated:
27
+ if not is_loopback(config.host):
28
+ log.warning("serving WITHOUT authentication on non-loopback address %s", config.host)
29
+ return
30
+
31
+ if not is_loopback(config.host):
32
+ raise InsecureConfiguration(
33
+ f"refusing to listen on {config.host} without API keys; configure security.api_keys / "
34
+ f"api_keys_file / ${config.security.api_keys_env}, or set security.allow_unauthenticated"
35
+ )
36
+
37
+ config.security.allow_unauthenticated = True
38
+ keys.allow_unauthenticated = True
39
+ log.warning("no API keys configured; accepting unauthenticated requests on loopback %s", config.host)
40
+
41
+
42
+ class APIServer:
43
+
44
+ def __init__(self, config):
45
+ self.config = config.validate()
46
+ self.models = ModelServer(self.config)
47
+ ui_dir = self.config.ui.directory or DEFAULT_UI_DIR
48
+ self.app = APIApp(self.models, self.config, ui_directory=ui_dir)
49
+ check_security(self.config, self.app.keys)
50
+
51
+ cors = {"Server": "pytensorforge"}
52
+ self.http = HTTPServer(self.app, self.config.host, self.config.port, self.config.limits,
53
+ on_request_done=self.app.on_request_done, extra_headers=cors)
54
+ self.http.on_connection_change = lambda n: self.app.metrics.open_connections.set(n)
55
+ self._loop = None
56
+ self._thread = None
57
+ self._stopped = None
58
+ self._stop_requested = None
59
+
60
+ @property
61
+ def port(self):
62
+ return self.http.bound_port
63
+
64
+ async def _run(self, ready=None, install_signals=True):
65
+ self._loop = asyncio.get_running_loop()
66
+ self._stop_requested = asyncio.Event()
67
+
68
+ await self._loop.run_in_executor(None, self.models.preload)
69
+ await self.http.start()
70
+ log.info("listening on http://%s:%d", self.config.host, self.port)
71
+
72
+ if install_signals:
73
+ for sig in (signal.SIGINT, signal.SIGTERM):
74
+ try:
75
+ self._loop.add_signal_handler(sig, self._stop_requested.set)
76
+ except (NotImplementedError, RuntimeError):
77
+ pass
78
+
79
+ if ready is not None:
80
+ ready.set()
81
+
82
+ await self._stop_requested.wait()
83
+ log.info("shutting down: draining in-flight requests (up to %.0fs)", self.config.limits.shutdown_drain_s)
84
+
85
+ await self.http.shutdown(self.config.limits.shutdown_drain_s)
86
+ await self._loop.run_in_executor(None, self.models.shutdown, 5.0)
87
+ self.app.close()
88
+ log.info("shutdown complete")
89
+
90
+ def serve_forever(self):
91
+ asyncio.run(self._run())
92
+
93
+ def start_background(self, timeout=60.0):
94
+ ready = threading.Event()
95
+ failure = []
96
+
97
+ def target():
98
+ try:
99
+ asyncio.run(self._run(ready=ready, install_signals=False))
100
+ except BaseException as exc:
101
+ failure.append(exc)
102
+ ready.set()
103
+
104
+ self._thread = threading.Thread(target=target, name="ptf-api", daemon=True)
105
+ self._thread.start()
106
+
107
+ if not ready.wait(timeout):
108
+ raise TimeoutError("server did not start in time")
109
+
110
+ if failure:
111
+ raise failure[0]
112
+
113
+ return self
114
+
115
+ def stop(self, timeout=30.0):
116
+ if self._loop is not None and self._stop_requested is not None:
117
+ self._loop.call_soon_threadsafe(self._stop_requested.set)
118
+
119
+ if self._thread is not None:
120
+ self._thread.join(timeout)
121
+ self._thread = None
File without changes