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.
- cli.py +604 -0
- pytensorforge-0.1.0.dist-info/METADATA +103 -0
- pytensorforge-0.1.0.dist-info/RECORD +146 -0
- pytensorforge-0.1.0.dist-info/WHEEL +5 -0
- pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
- pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
- src/__init__.py +0 -0
- src/activations/Activation.py +4 -0
- src/activations/ELU.py +11 -0
- src/activations/GELU.py +6 -0
- src/activations/ReLU.py +27 -0
- src/activations/SELU.py +14 -0
- src/activations/Sigmoid.py +27 -0
- src/activations/Softmax.py +84 -0
- src/activations/Tanh.py +29 -0
- src/activations/__init__.py +17 -0
- src/config.py +120 -0
- src/core/Matrix.py +3 -0
- src/core/Scalar.py +18 -0
- src/core/Tensor.py +866 -0
- src/core/Vector.py +31 -0
- src/core/__init__.py +0 -0
- src/data/__init__.py +0 -0
- src/data/chat_dataset.py +188 -0
- src/data/corpus.py +104 -0
- src/data/document_stream.py +178 -0
- src/data/parallel_encode.py +86 -0
- src/data/prefetch.py +62 -0
- src/data/shard_builder.py +119 -0
- src/data/shard_writer.py +81 -0
- src/data/sharded_dataset.py +112 -0
- src/data/streaming_dataset.py +132 -0
- src/data/validation.py +212 -0
- src/inference/__init__.py +0 -0
- src/inference/chat_template.py +384 -0
- src/inference/config.py +48 -0
- src/inference/engine.py +241 -0
- src/inference/export.py +133 -0
- src/inference/kv_cache.py +65 -0
- src/inference/runtime.py +161 -0
- src/inference/sampling.py +42 -0
- src/inference/scheduler.py +473 -0
- src/inference/text.py +67 -0
- src/initializers/Constant.py +9 -0
- src/initializers/GlorotNormal.py +15 -0
- src/initializers/GlorotUniform.py +26 -0
- src/initializers/HeNormal.py +15 -0
- src/initializers/HeUniform.py +14 -0
- src/initializers/Initializer.py +4 -0
- src/initializers/LecunNormal.py +16 -0
- src/initializers/LecunUniform.py +14 -0
- src/initializers/Ones.py +6 -0
- src/initializers/Orthogonal.py +14 -0
- src/initializers/RandomNormal.py +14 -0
- src/initializers/RandomUniform.py +14 -0
- src/initializers/Zeros.py +8 -0
- src/initializers/__init__.py +17 -0
- src/loss/CategoricalCrossEntropy.py +9 -0
- src/loss/CrossEntropyLoss.py +34 -0
- src/loss/CrossEntropyWithLogitsLoss.py +59 -0
- src/loss/Hinge.py +5 -0
- src/loss/Huber.py +22 -0
- src/loss/Loss.py +6 -0
- src/loss/MSE.py +7 -0
- src/loss/MSELoss.py +10 -0
- src/loss/SparseCategoricalCrossEntropy.py +15 -0
- src/loss/__init__.py +18 -0
- src/loss/bce.py +34 -0
- src/loss/mae.py +16 -0
- src/math/__init__.py +0 -0
- src/math/clip.py +37 -0
- src/math/exp.py +27 -0
- src/math/log.py +25 -0
- src/math/sigmoid.py +5 -0
- src/models/__init__.py +0 -0
- src/models/embedding/Embedding.py +65 -0
- src/models/embedding/__init__.py +0 -0
- src/models/gpt/__init__.py +0 -0
- src/models/gpt/attention.py +158 -0
- src/models/gpt/block.py +74 -0
- src/models/gpt/config.py +103 -0
- src/models/gpt/context.py +44 -0
- src/models/gpt/model.py +165 -0
- src/models/gpt/recompute.py +35 -0
- src/models/gpt/rope.py +84 -0
- src/models/regression/Linear.py +51 -0
- src/models/regression/Logistic.py +36 -0
- src/models/regression/__init__.py +0 -0
- src/models/seq/Sequential.py +297 -0
- src/models/seq/__init__.py +0 -0
- src/models/svm/__init__.py +0 -0
- src/models/tokenizer/BPETokenizer.py +228 -0
- src/models/tokenizer/__init__.py +0 -0
- src/models/transformers/Dropout.py +35 -0
- src/models/transformers/LastToken.py +10 -0
- src/models/transformers/LayerNorm.py +54 -0
- src/models/transformers/Linear.py +18 -0
- src/models/transformers/MultiHeadAttention.py +130 -0
- src/models/transformers/TransformerBlock.py +79 -0
- src/models/transformers/__init__.py +0 -0
- src/neural/Dense.py +58 -0
- src/neural/LSTM.py +167 -0
- src/neural/Layer.py +72 -0
- src/neural/Parameter.py +30 -0
- src/neural/RNN.py +83 -0
- src/neural/__init__.py +0 -0
- src/ops/__init__.py +0 -0
- src/ops/stack.py +40 -0
- src/optimizers/Adagrad.py +31 -0
- src/optimizers/Adam.py +98 -0
- src/optimizers/AdamW.py +84 -0
- src/optimizers/Batch.py +11 -0
- src/optimizers/Nesterov.py +35 -0
- src/optimizers/Optimizer.py +18 -0
- src/optimizers/RMSProp.py +35 -0
- src/optimizers/SGD.py +30 -0
- src/optimizers/SGDMomentum.py +28 -0
- src/optimizers/__init__.py +9 -0
- src/scaling/StandardScaler.py +15 -0
- src/scaling/__init__.py +0 -0
- src/serialization/__init__.py +0 -0
- src/serialization/checkpoint.py +58 -0
- src/serialization/modelio.py +132 -0
- src/serving/__init__.py +0 -0
- src/serving/app.py +792 -0
- src/serving/config.py +216 -0
- src/serving/errors.py +51 -0
- src/serving/http.py +599 -0
- src/serving/metrics.py +293 -0
- src/serving/model_server.py +287 -0
- src/serving/protocol.py +377 -0
- src/serving/security.py +200 -0
- src/serving/server.py +121 -0
- src/tokenization/__init__.py +0 -0
- src/tokenization/base.py +75 -0
- src/tokenization/bpe.py +190 -0
- src/tokenization/bytebpe.py +476 -0
- src/tokenization/registry.py +28 -0
- src/training/__init__.py +0 -0
- src/training/checkpoint_manager.py +101 -0
- src/training/experiment.py +71 -0
- src/training/losses.py +42 -0
- src/training/precision.py +141 -0
- src/training/profiler.py +38 -0
- src/training/scheduler.py +50 -0
- src/training/trainer.py +594 -0
src/serving/http.py
ADDED
|
@@ -0,0 +1,599 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import json
|
|
3
|
+
import logging
|
|
4
|
+
import time
|
|
5
|
+
import uuid
|
|
6
|
+
from http import HTTPStatus
|
|
7
|
+
from urllib.parse import parse_qsl, unquote, urlsplit
|
|
8
|
+
|
|
9
|
+
log = logging.getLogger("ptf.http")
|
|
10
|
+
|
|
11
|
+
MAX_REQUEST_LINE = 8192
|
|
12
|
+
READ_CHUNK = 65536
|
|
13
|
+
TOKEN_CHARS = set("!#$%&'*+-.^_`|~0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class HTTPError(Exception):
|
|
17
|
+
|
|
18
|
+
def __init__(self, status, message, close=True):
|
|
19
|
+
super().__init__(message)
|
|
20
|
+
self.status = status
|
|
21
|
+
self.message = message
|
|
22
|
+
self.close = close
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class ClientDisconnected(Exception):
|
|
26
|
+
pass
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class Request:
|
|
30
|
+
__slots__ = ("method", "target", "path", "query", "version", "headers", "body", "client_ip", "request_id",
|
|
31
|
+
"received_at", "info")
|
|
32
|
+
|
|
33
|
+
def __init__(self, method, target, version, headers, client_ip):
|
|
34
|
+
self.method = method
|
|
35
|
+
self.target = target
|
|
36
|
+
parts = urlsplit(target)
|
|
37
|
+
self.path = unquote(parts.path) or "/"
|
|
38
|
+
self.query = dict(parse_qsl(parts.query, keep_blank_values=True))
|
|
39
|
+
self.version = version
|
|
40
|
+
self.headers = headers
|
|
41
|
+
self.body = b""
|
|
42
|
+
self.client_ip = client_ip
|
|
43
|
+
self.request_id = uuid.uuid4().hex[:16]
|
|
44
|
+
self.received_at = time.perf_counter()
|
|
45
|
+
self.info = {}
|
|
46
|
+
|
|
47
|
+
def json(self):
|
|
48
|
+
if not self.body:
|
|
49
|
+
raise ValueError("empty body")
|
|
50
|
+
|
|
51
|
+
return json.loads(self.body)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class Response:
|
|
55
|
+
|
|
56
|
+
def __init__(self, status=200, body=b"", headers=None, content_type="application/json"):
|
|
57
|
+
self.status = status
|
|
58
|
+
self.headers = dict(headers or {})
|
|
59
|
+
|
|
60
|
+
if isinstance(body, (dict, list)):
|
|
61
|
+
body = json.dumps(body, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
|
|
62
|
+
elif isinstance(body, str):
|
|
63
|
+
body = body.encode("utf-8")
|
|
64
|
+
|
|
65
|
+
self.body = body
|
|
66
|
+
|
|
67
|
+
if content_type and "Content-Type" not in self.headers:
|
|
68
|
+
self.headers["Content-Type"] = content_type
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class StreamingResponse:
|
|
72
|
+
|
|
73
|
+
def __init__(self, iterator, status=200, headers=None, content_type="text/event-stream; charset=utf-8",
|
|
74
|
+
on_close=None):
|
|
75
|
+
self.iterator = iterator
|
|
76
|
+
self.on_close = on_close
|
|
77
|
+
self.status = status
|
|
78
|
+
self.headers = dict(headers or {})
|
|
79
|
+
self.headers.setdefault("Content-Type", content_type)
|
|
80
|
+
self.headers.setdefault("Cache-Control", "no-cache")
|
|
81
|
+
self.headers.setdefault("X-Accel-Buffering", "no")
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
class ConnectionContext:
|
|
85
|
+
|
|
86
|
+
def __init__(self, conn):
|
|
87
|
+
self._conn = conn
|
|
88
|
+
self.disconnected = asyncio.Event()
|
|
89
|
+
self.write_seconds = 0.0
|
|
90
|
+
|
|
91
|
+
@property
|
|
92
|
+
def client_ip(self):
|
|
93
|
+
return self._conn.client_ip
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class _Connection:
|
|
97
|
+
|
|
98
|
+
def __init__(self, reader, writer, limits):
|
|
99
|
+
self.reader = reader
|
|
100
|
+
self.writer = writer
|
|
101
|
+
self.limits = limits
|
|
102
|
+
self.buf = bytearray()
|
|
103
|
+
self.eof = False
|
|
104
|
+
peer = writer.get_extra_info("peername")
|
|
105
|
+
self.client_ip = peer[0] if isinstance(peer, tuple) else "unknown"
|
|
106
|
+
|
|
107
|
+
async def _fill(self, timeout):
|
|
108
|
+
if self.eof:
|
|
109
|
+
return False
|
|
110
|
+
|
|
111
|
+
if timeout is None:
|
|
112
|
+
data = await self.reader.read(READ_CHUNK)
|
|
113
|
+
else:
|
|
114
|
+
data = await asyncio.wait_for(self.reader.read(READ_CHUNK), timeout)
|
|
115
|
+
|
|
116
|
+
if not data:
|
|
117
|
+
self.eof = True
|
|
118
|
+
return False
|
|
119
|
+
|
|
120
|
+
self.buf += data
|
|
121
|
+
return True
|
|
122
|
+
|
|
123
|
+
async def read_head(self, first_timeout, rest_timeout):
|
|
124
|
+
timeout = first_timeout
|
|
125
|
+
deadline = None
|
|
126
|
+
|
|
127
|
+
while True:
|
|
128
|
+
idx = self.buf.find(b"\r\n\r\n")
|
|
129
|
+
|
|
130
|
+
if idx != -1:
|
|
131
|
+
if idx + 4 > self.limits.max_header_bytes:
|
|
132
|
+
raise HTTPError(431, "request headers too large")
|
|
133
|
+
|
|
134
|
+
head = bytes(self.buf[:idx])
|
|
135
|
+
del self.buf[:idx + 4]
|
|
136
|
+
return head
|
|
137
|
+
|
|
138
|
+
if len(self.buf) > self.limits.max_header_bytes:
|
|
139
|
+
raise HTTPError(431, "request headers too large")
|
|
140
|
+
|
|
141
|
+
if self.buf and deadline is None:
|
|
142
|
+
deadline = time.monotonic() + rest_timeout
|
|
143
|
+
|
|
144
|
+
if deadline is not None:
|
|
145
|
+
timeout = max(0.001, deadline - time.monotonic())
|
|
146
|
+
|
|
147
|
+
try:
|
|
148
|
+
got = await self._fill(timeout)
|
|
149
|
+
except asyncio.TimeoutError:
|
|
150
|
+
if self.buf:
|
|
151
|
+
raise HTTPError(408, "timed out reading request headers")
|
|
152
|
+
return None
|
|
153
|
+
|
|
154
|
+
if not got:
|
|
155
|
+
if self.buf:
|
|
156
|
+
raise ClientDisconnected()
|
|
157
|
+
return None
|
|
158
|
+
|
|
159
|
+
async def read_exactly(self, n, timeout):
|
|
160
|
+
deadline = time.monotonic() + timeout
|
|
161
|
+
|
|
162
|
+
while len(self.buf) < n:
|
|
163
|
+
left = deadline - time.monotonic()
|
|
164
|
+
|
|
165
|
+
if left <= 0:
|
|
166
|
+
raise HTTPError(408, "timed out reading request body")
|
|
167
|
+
|
|
168
|
+
try:
|
|
169
|
+
got = await self._fill(left)
|
|
170
|
+
except asyncio.TimeoutError:
|
|
171
|
+
raise HTTPError(408, "timed out reading request body")
|
|
172
|
+
|
|
173
|
+
if not got:
|
|
174
|
+
raise ClientDisconnected()
|
|
175
|
+
|
|
176
|
+
data = bytes(self.buf[:n])
|
|
177
|
+
del self.buf[:n]
|
|
178
|
+
return data
|
|
179
|
+
|
|
180
|
+
async def linger(self, timeout, max_bytes):
|
|
181
|
+
try:
|
|
182
|
+
self.writer.write_eof()
|
|
183
|
+
except (OSError, RuntimeError, AttributeError):
|
|
184
|
+
pass
|
|
185
|
+
|
|
186
|
+
deadline = time.monotonic() + timeout
|
|
187
|
+
discarded = 0
|
|
188
|
+
|
|
189
|
+
while discarded < max_bytes and not self.eof:
|
|
190
|
+
left = deadline - time.monotonic()
|
|
191
|
+
|
|
192
|
+
if left <= 0:
|
|
193
|
+
return
|
|
194
|
+
|
|
195
|
+
self.buf.clear()
|
|
196
|
+
|
|
197
|
+
try:
|
|
198
|
+
data = await asyncio.wait_for(self.reader.read(READ_CHUNK), left)
|
|
199
|
+
except (asyncio.TimeoutError, ConnectionError, OSError):
|
|
200
|
+
return
|
|
201
|
+
|
|
202
|
+
if not data:
|
|
203
|
+
return
|
|
204
|
+
|
|
205
|
+
discarded += len(data)
|
|
206
|
+
|
|
207
|
+
async def watch_disconnect(self, ctx):
|
|
208
|
+
cap = self.limits.max_header_bytes + self.limits.max_body_bytes
|
|
209
|
+
|
|
210
|
+
try:
|
|
211
|
+
while not self.eof and len(self.buf) <= cap:
|
|
212
|
+
await self._fill(None)
|
|
213
|
+
except (ConnectionError, OSError):
|
|
214
|
+
self.eof = True
|
|
215
|
+
|
|
216
|
+
if self.eof:
|
|
217
|
+
ctx.disconnected.set()
|
|
218
|
+
|
|
219
|
+
async def write(self, data, ctx):
|
|
220
|
+
if self.writer.is_closing():
|
|
221
|
+
raise ClientDisconnected()
|
|
222
|
+
|
|
223
|
+
self.writer.write(data)
|
|
224
|
+
started = time.perf_counter()
|
|
225
|
+
|
|
226
|
+
try:
|
|
227
|
+
await self.writer.drain()
|
|
228
|
+
except (ConnectionError, OSError) as exc:
|
|
229
|
+
raise ClientDisconnected() from exc
|
|
230
|
+
finally:
|
|
231
|
+
if ctx is not None:
|
|
232
|
+
ctx.write_seconds += time.perf_counter() - started
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
def _parse_head(head, client_ip):
|
|
236
|
+
try:
|
|
237
|
+
text = head.decode("latin-1")
|
|
238
|
+
except UnicodeDecodeError:
|
|
239
|
+
raise HTTPError(400, "malformed request")
|
|
240
|
+
|
|
241
|
+
lines = text.split("\r\n")
|
|
242
|
+
request_line = lines[0]
|
|
243
|
+
|
|
244
|
+
if len(request_line) > MAX_REQUEST_LINE:
|
|
245
|
+
raise HTTPError(414, "request line too long")
|
|
246
|
+
|
|
247
|
+
parts = request_line.split(" ")
|
|
248
|
+
|
|
249
|
+
if len(parts) != 3:
|
|
250
|
+
raise HTTPError(400, "malformed request line")
|
|
251
|
+
|
|
252
|
+
method, target, version = parts
|
|
253
|
+
|
|
254
|
+
if version not in ("HTTP/1.1", "HTTP/1.0"):
|
|
255
|
+
raise HTTPError(505, "HTTP version not supported")
|
|
256
|
+
|
|
257
|
+
if not method or any(c not in TOKEN_CHARS for c in method):
|
|
258
|
+
raise HTTPError(400, "malformed method")
|
|
259
|
+
|
|
260
|
+
if not target.startswith("/"):
|
|
261
|
+
raise HTTPError(400, "only origin-form request targets are supported")
|
|
262
|
+
|
|
263
|
+
headers = {}
|
|
264
|
+
|
|
265
|
+
for line in lines[1:]:
|
|
266
|
+
if not line:
|
|
267
|
+
continue
|
|
268
|
+
|
|
269
|
+
if line[0] in " \t":
|
|
270
|
+
raise HTTPError(400, "obsolete header folding is not allowed")
|
|
271
|
+
|
|
272
|
+
name, sep, value = line.partition(":")
|
|
273
|
+
|
|
274
|
+
if not sep or not name or any(c not in TOKEN_CHARS for c in name):
|
|
275
|
+
raise HTTPError(400, "malformed header")
|
|
276
|
+
|
|
277
|
+
name = name.lower()
|
|
278
|
+
value = value.strip()
|
|
279
|
+
|
|
280
|
+
if name in headers:
|
|
281
|
+
if name in ("content-length", "host", "authorization", "transfer-encoding"):
|
|
282
|
+
if headers[name] != value:
|
|
283
|
+
raise HTTPError(400, f"conflicting duplicate '{name}' header")
|
|
284
|
+
continue
|
|
285
|
+
|
|
286
|
+
headers[name] = headers[name] + ", " + value
|
|
287
|
+
else:
|
|
288
|
+
headers[name] = value
|
|
289
|
+
|
|
290
|
+
return Request(method, target, version, headers, client_ip)
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def _status_line(status):
|
|
294
|
+
try:
|
|
295
|
+
phrase = HTTPStatus(status).phrase
|
|
296
|
+
except ValueError:
|
|
297
|
+
phrase = "Unknown"
|
|
298
|
+
|
|
299
|
+
return f"HTTP/1.1 {status} {phrase}\r\n"
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
def _encode_headers(status, headers):
|
|
303
|
+
out = [_status_line(status)]
|
|
304
|
+
|
|
305
|
+
for k, v in headers.items():
|
|
306
|
+
v = str(v)
|
|
307
|
+
|
|
308
|
+
if "\r" in v or "\n" in v or "\r" in k or "\n" in k:
|
|
309
|
+
raise ValueError("header injection attempt")
|
|
310
|
+
|
|
311
|
+
out.append(f"{k}: {v}\r\n")
|
|
312
|
+
|
|
313
|
+
out.append("\r\n")
|
|
314
|
+
return "".join(out).encode("latin-1")
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
def _error_body(status, message):
|
|
318
|
+
kind = "invalid_request_error" if status < 500 else "server_error"
|
|
319
|
+
return json.dumps({"error": {"message": message, "type": kind, "param": None, "code": None}}).encode()
|
|
320
|
+
|
|
321
|
+
|
|
322
|
+
class HTTPServer:
|
|
323
|
+
|
|
324
|
+
def __init__(self, handler, host, port, limits, on_request_done=None, extra_headers=None):
|
|
325
|
+
self.handler = handler
|
|
326
|
+
self.host = host
|
|
327
|
+
self.port = port
|
|
328
|
+
self.limits = limits
|
|
329
|
+
self.on_request_done = on_request_done
|
|
330
|
+
self.extra_headers = dict(extra_headers or {})
|
|
331
|
+
self._server = None
|
|
332
|
+
self._connections = set()
|
|
333
|
+
self._busy = set()
|
|
334
|
+
self._active_requests = 0
|
|
335
|
+
self._closing = False
|
|
336
|
+
self._idle = asyncio.Event()
|
|
337
|
+
self._idle.set()
|
|
338
|
+
self.on_connection_change = None
|
|
339
|
+
|
|
340
|
+
@property
|
|
341
|
+
def sockets(self):
|
|
342
|
+
return self._server.sockets if self._server else []
|
|
343
|
+
|
|
344
|
+
@property
|
|
345
|
+
def bound_port(self):
|
|
346
|
+
socks = self.sockets
|
|
347
|
+
return socks[0].getsockname()[1] if socks else None
|
|
348
|
+
|
|
349
|
+
async def start(self):
|
|
350
|
+
self._server = await asyncio.start_server(
|
|
351
|
+
self._on_connection, self.host, self.port, limit=READ_CHUNK * 4, backlog=1024,
|
|
352
|
+
)
|
|
353
|
+
|
|
354
|
+
def _close_idle(self):
|
|
355
|
+
for writer in list(self._connections):
|
|
356
|
+
if writer not in self._busy:
|
|
357
|
+
writer.close()
|
|
358
|
+
|
|
359
|
+
async def shutdown(self, drain_timeout):
|
|
360
|
+
self._closing = True
|
|
361
|
+
|
|
362
|
+
if self._server is not None:
|
|
363
|
+
self._server.close()
|
|
364
|
+
|
|
365
|
+
self._close_idle()
|
|
366
|
+
|
|
367
|
+
try:
|
|
368
|
+
await asyncio.wait_for(self._idle.wait(), drain_timeout)
|
|
369
|
+
except asyncio.TimeoutError:
|
|
370
|
+
log.warning("shutdown drain timed out with %d active requests", self._active_requests)
|
|
371
|
+
|
|
372
|
+
for writer in list(self._connections):
|
|
373
|
+
writer.close()
|
|
374
|
+
|
|
375
|
+
if self._server is not None:
|
|
376
|
+
try:
|
|
377
|
+
await asyncio.wait_for(self._server.wait_closed(), 5.0)
|
|
378
|
+
except asyncio.TimeoutError:
|
|
379
|
+
log.warning("some connections did not close cleanly")
|
|
380
|
+
|
|
381
|
+
def _conn_changed(self):
|
|
382
|
+
if self.on_connection_change:
|
|
383
|
+
self.on_connection_change(len(self._connections))
|
|
384
|
+
|
|
385
|
+
async def _on_connection(self, reader, writer):
|
|
386
|
+
if len(self._connections) >= self.limits.max_connections or self._closing:
|
|
387
|
+
try:
|
|
388
|
+
body = _error_body(503, "too many connections")
|
|
389
|
+
writer.write(_encode_headers(503, {
|
|
390
|
+
"Content-Type": "application/json", "Content-Length": len(body),
|
|
391
|
+
"Connection": "close", "Retry-After": "1",
|
|
392
|
+
}) + body)
|
|
393
|
+
await asyncio.wait_for(writer.drain(), 1.0)
|
|
394
|
+
except Exception:
|
|
395
|
+
pass
|
|
396
|
+
finally:
|
|
397
|
+
writer.close()
|
|
398
|
+
return
|
|
399
|
+
|
|
400
|
+
self._connections.add(writer)
|
|
401
|
+
self._conn_changed()
|
|
402
|
+
conn = _Connection(reader, writer, self.limits)
|
|
403
|
+
|
|
404
|
+
try:
|
|
405
|
+
await self._serve(conn)
|
|
406
|
+
except (ClientDisconnected, ConnectionError, OSError, asyncio.IncompleteReadError):
|
|
407
|
+
pass
|
|
408
|
+
except asyncio.CancelledError:
|
|
409
|
+
raise
|
|
410
|
+
except Exception:
|
|
411
|
+
log.exception("unhandled connection error")
|
|
412
|
+
finally:
|
|
413
|
+
self._connections.discard(writer)
|
|
414
|
+
self._busy.discard(writer)
|
|
415
|
+
self._conn_changed()
|
|
416
|
+
|
|
417
|
+
try:
|
|
418
|
+
writer.close()
|
|
419
|
+
except Exception:
|
|
420
|
+
pass
|
|
421
|
+
|
|
422
|
+
async def _serve(self, conn):
|
|
423
|
+
served = 0
|
|
424
|
+
|
|
425
|
+
while not self._closing:
|
|
426
|
+
first_timeout = self.limits.header_timeout_s if served == 0 else self.limits.keepalive_timeout_s
|
|
427
|
+
|
|
428
|
+
try:
|
|
429
|
+
head = await conn.read_head(first_timeout, self.limits.header_timeout_s)
|
|
430
|
+
except HTTPError as exc:
|
|
431
|
+
await self._send_simple(conn, exc.status, exc.message, keep_alive=False)
|
|
432
|
+
return
|
|
433
|
+
|
|
434
|
+
if head is None:
|
|
435
|
+
return
|
|
436
|
+
|
|
437
|
+
try:
|
|
438
|
+
request = _parse_head(head, conn.client_ip)
|
|
439
|
+
keep_alive = self._wants_keep_alive(request)
|
|
440
|
+
await self._read_body(conn, request)
|
|
441
|
+
except HTTPError as exc:
|
|
442
|
+
await self._send_simple(conn, exc.status, exc.message, keep_alive=False)
|
|
443
|
+
|
|
444
|
+
if exc.status == 413:
|
|
445
|
+
await conn.linger(2.0, 4 * self.limits.max_body_bytes)
|
|
446
|
+
|
|
447
|
+
return
|
|
448
|
+
|
|
449
|
+
served += 1
|
|
450
|
+
|
|
451
|
+
if served >= self.limits.max_requests_per_connection:
|
|
452
|
+
keep_alive = False
|
|
453
|
+
|
|
454
|
+
keep_alive = await self._dispatch(conn, request, keep_alive and not self._closing)
|
|
455
|
+
|
|
456
|
+
if not keep_alive:
|
|
457
|
+
return
|
|
458
|
+
|
|
459
|
+
def _wants_keep_alive(self, request):
|
|
460
|
+
conn_header = request.headers.get("connection", "").lower()
|
|
461
|
+
|
|
462
|
+
if request.version == "HTTP/1.0":
|
|
463
|
+
return "keep-alive" in conn_header
|
|
464
|
+
|
|
465
|
+
return "close" not in conn_header
|
|
466
|
+
|
|
467
|
+
async def _read_body(self, conn, request):
|
|
468
|
+
if "transfer-encoding" in request.headers:
|
|
469
|
+
raise HTTPError(501, "chunked request bodies are not supported; send Content-Length")
|
|
470
|
+
|
|
471
|
+
raw = request.headers.get("content-length")
|
|
472
|
+
|
|
473
|
+
if raw is None:
|
|
474
|
+
if request.method in ("POST", "PUT", "PATCH"):
|
|
475
|
+
raise HTTPError(411, "Content-Length is required")
|
|
476
|
+
return
|
|
477
|
+
|
|
478
|
+
if not raw.isdigit():
|
|
479
|
+
raise HTTPError(400, "invalid Content-Length")
|
|
480
|
+
|
|
481
|
+
length = int(raw)
|
|
482
|
+
|
|
483
|
+
if length > self.limits.max_body_bytes:
|
|
484
|
+
raise HTTPError(413, f"request body exceeds {self.limits.max_body_bytes} bytes")
|
|
485
|
+
|
|
486
|
+
if request.headers.get("expect", "").lower() == "100-continue" and length > len(conn.buf):
|
|
487
|
+
await conn.write(b"HTTP/1.1 100 Continue\r\n\r\n", None)
|
|
488
|
+
|
|
489
|
+
request.body = await conn.read_exactly(length, self.limits.body_timeout_s)
|
|
490
|
+
|
|
491
|
+
async def _send_simple(self, conn, status, message, keep_alive):
|
|
492
|
+
body = _error_body(status, message)
|
|
493
|
+
headers = {"Content-Type": "application/json", "Content-Length": len(body),
|
|
494
|
+
"Connection": "keep-alive" if keep_alive else "close"}
|
|
495
|
+
headers.update(self.extra_headers)
|
|
496
|
+
|
|
497
|
+
try:
|
|
498
|
+
await conn.write(_encode_headers(status, headers) + body, None)
|
|
499
|
+
except ClientDisconnected:
|
|
500
|
+
pass
|
|
501
|
+
|
|
502
|
+
async def _dispatch(self, conn, request, keep_alive):
|
|
503
|
+
ctx = ConnectionContext(conn)
|
|
504
|
+
self._busy.add(conn.writer)
|
|
505
|
+
watcher = asyncio.ensure_future(conn.watch_disconnect(ctx))
|
|
506
|
+
self._active_requests += 1
|
|
507
|
+
self._idle.clear()
|
|
508
|
+
status = 500
|
|
509
|
+
streamed = False
|
|
510
|
+
|
|
511
|
+
try:
|
|
512
|
+
try:
|
|
513
|
+
response = await self.handler(request, ctx)
|
|
514
|
+
except ClientDisconnected:
|
|
515
|
+
return False
|
|
516
|
+
|
|
517
|
+
status = response.status
|
|
518
|
+
headers = dict(self.extra_headers)
|
|
519
|
+
headers.update(response.headers)
|
|
520
|
+
headers["X-Request-Id"] = request.request_id
|
|
521
|
+
|
|
522
|
+
if isinstance(response, StreamingResponse):
|
|
523
|
+
streamed = True
|
|
524
|
+
return await self._write_stream(conn, request, response, headers, ctx, keep_alive)
|
|
525
|
+
|
|
526
|
+
body = b"" if request.method == "HEAD" else response.body
|
|
527
|
+
headers["Content-Length"] = len(response.body)
|
|
528
|
+
headers["Connection"] = "keep-alive" if keep_alive else "close"
|
|
529
|
+
await conn.write(_encode_headers(response.status, headers) + body, ctx)
|
|
530
|
+
|
|
531
|
+
return keep_alive and not conn.eof
|
|
532
|
+
except ClientDisconnected:
|
|
533
|
+
return False
|
|
534
|
+
finally:
|
|
535
|
+
if not watcher.done():
|
|
536
|
+
watcher.cancel()
|
|
537
|
+
|
|
538
|
+
try:
|
|
539
|
+
await watcher
|
|
540
|
+
except (asyncio.CancelledError, Exception):
|
|
541
|
+
pass
|
|
542
|
+
|
|
543
|
+
self._active_requests -= 1
|
|
544
|
+
self._busy.discard(conn.writer)
|
|
545
|
+
|
|
546
|
+
if self._active_requests == 0:
|
|
547
|
+
self._idle.set()
|
|
548
|
+
|
|
549
|
+
if self.on_request_done is not None:
|
|
550
|
+
try:
|
|
551
|
+
self.on_request_done(request, status, time.perf_counter() - request.received_at,
|
|
552
|
+
ctx.write_seconds, ctx.disconnected.is_set(), streamed)
|
|
553
|
+
except Exception:
|
|
554
|
+
log.exception("request completion hook failed")
|
|
555
|
+
|
|
556
|
+
async def _write_stream(self, conn, request, response, headers, ctx, keep_alive):
|
|
557
|
+
chunked = request.version == "HTTP/1.1"
|
|
558
|
+
|
|
559
|
+
if chunked:
|
|
560
|
+
headers["Transfer-Encoding"] = "chunked"
|
|
561
|
+
headers["Connection"] = "keep-alive" if keep_alive else "close"
|
|
562
|
+
else:
|
|
563
|
+
headers["Connection"] = "close"
|
|
564
|
+
keep_alive = False
|
|
565
|
+
|
|
566
|
+
iterator = response.iterator
|
|
567
|
+
completed = False
|
|
568
|
+
|
|
569
|
+
try:
|
|
570
|
+
await conn.write(_encode_headers(response.status, headers), ctx)
|
|
571
|
+
|
|
572
|
+
async for piece in iterator:
|
|
573
|
+
if ctx.disconnected.is_set():
|
|
574
|
+
raise ClientDisconnected()
|
|
575
|
+
|
|
576
|
+
if not piece:
|
|
577
|
+
continue
|
|
578
|
+
|
|
579
|
+
if isinstance(piece, str):
|
|
580
|
+
piece = piece.encode("utf-8")
|
|
581
|
+
|
|
582
|
+
frame = b"%x\r\n%s\r\n" % (len(piece), piece) if chunked else piece
|
|
583
|
+
await conn.write(frame, ctx)
|
|
584
|
+
|
|
585
|
+
if chunked:
|
|
586
|
+
await conn.write(b"0\r\n\r\n", ctx)
|
|
587
|
+
|
|
588
|
+
completed = True
|
|
589
|
+
finally:
|
|
590
|
+
if not completed:
|
|
591
|
+
ctx.disconnected.set()
|
|
592
|
+
|
|
593
|
+
try:
|
|
594
|
+
await iterator.aclose()
|
|
595
|
+
finally:
|
|
596
|
+
if response.on_close is not None:
|
|
597
|
+
response.on_close()
|
|
598
|
+
|
|
599
|
+
return keep_alive and not conn.eof
|