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
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