Skip to content

Commit 0396151

Browse files
Refactor connection
1 parent 8adc636 commit 0396151

12 files changed

Lines changed: 364 additions & 219 deletions

File tree

pyrogram/__init__.py

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -38,5 +38,3 @@ class ContinuePropagation(StopAsyncIteration):
3838
from . import raw, types, filters, handlers, enums
3939
from .client import Client
4040
from .sync import idle, compose
41-
42-
crypto_executor = ThreadPoolExecutor(1, thread_name_prefix="CryptoWorker")

pyrogram/connection/connection.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ def __init__(
3737
proxy: dict,
3838
media: bool = False,
3939
protocol_factory: Type[TCP] = TCPAbridged,
40+
crypto_executor_workers: int = 1,
4041
loop: Optional[asyncio.AbstractEventLoop] = None
4142
) -> None:
4243
self.dc_id = dc_id
@@ -45,6 +46,7 @@ def __init__(
4546
self.proxy = proxy
4647
self.media = media
4748
self.protocol_factory = protocol_factory
49+
self.crypto_executor_workers = crypto_executor_workers
4850

4951
self.address = DataCenter(dc_id, test_mode, ipv6, media)
5052
self.protocol: Optional[TCP] = None
@@ -56,7 +58,7 @@ def __init__(
5658

5759
async def connect(self) -> None:
5860
for i in range(Connection.MAX_CONNECTION_ATTEMPTS):
59-
self.protocol = self.protocol_factory(ipv6=self.ipv6, proxy=self.proxy, loop=self.loop)
61+
self.protocol = self.protocol_factory(ipv6=self.ipv6, proxy=self.proxy, crypto_executor_workers=self.crypto_executor_workers, loop=self.loop)
6062

6163
try:
6264
log.info("Connecting...")

pyrogram/connection/transport/tcp/tcp.py

Lines changed: 40 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
import logging
2222
import socket
2323
from concurrent.futures import ThreadPoolExecutor
24-
from typing import Tuple, Dict, TypedDict, Optional
24+
from typing import Dict, Optional, Tuple, TypedDict
2525

2626
import socks
2727

@@ -45,13 +45,24 @@ class Proxy(TypedDict):
4545
class TCP:
4646
TIMEOUT = 10
4747

48-
def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEventLoop] = None) -> None:
48+
def __init__(
49+
self,
50+
ipv6: bool = False,
51+
proxy: Proxy = None,
52+
crypto_executor_workers: int = 1,
53+
loop: Optional[asyncio.AbstractEventLoop] = None,
54+
) -> None:
4955
self.ipv6 = ipv6
5056
self.proxy = proxy
5157

58+
self.crypto_executor_workers = crypto_executor_workers
59+
self.crypto_executor = ThreadPoolExecutor(max_workers=self.crypto_executor_workers, thread_name_prefix="CryptoWorker")
60+
5261
self.reader: Optional[asyncio.StreamReader] = None
5362
self.writer: Optional[asyncio.StreamWriter] = None
5463

64+
self.marker_event = asyncio.Event()
65+
5566
self.lock = asyncio.Lock()
5667

5768
if isinstance(loop, asyncio.AbstractEventLoop):
@@ -95,8 +106,7 @@ async def _connect_via_proxy(
95106
)
96107
sock.settimeout(TCP.TIMEOUT)
97108

98-
with ThreadPoolExecutor() as executor:
99-
await self.loop.run_in_executor(executor, sock.connect, destination)
109+
await self.loop.run_in_executor(self.crypto_executor, sock.connect, destination)
100110

101111
sock.setblocking(False)
102112

@@ -124,27 +134,40 @@ async def _connect(self, destination: Tuple[str, int]) -> None:
124134

125135
async def connect(self, address: Tuple[str, int]) -> None:
126136
try:
127-
await asyncio.wait_for(self._connect(address), TCP.TIMEOUT)
137+
await asyncio.wait_for(self._connect(address), timeout=TCP.TIMEOUT)
128138
except asyncio.TimeoutError: # Re-raise as TimeoutError. asyncio.TimeoutError is deprecated in 3.11
129139
raise TimeoutError("Connection timed out")
130140

131141
async def close(self) -> None:
132-
if self.writer is None:
133-
return None
142+
async with self.lock:
143+
if self.writer is None or self.writer.is_closing():
144+
return None
134145

135-
try:
136-
self.writer.close()
137-
await asyncio.wait_for(self.writer.wait_closed(), TCP.TIMEOUT)
138-
except Exception as e:
139-
log.info("Close exception: %s %s", type(e).__name__, e)
140-
finally:
141-
self.writer = None
142-
143-
async def send(self, data: bytes) -> None:
146+
try:
147+
if self.writer.transport is not None:
148+
self.writer.transport.abort()
149+
150+
self.writer.close()
151+
152+
await asyncio.wait_for(self.writer.wait_closed(), timeout=TCP.TIMEOUT)
153+
except asyncio.TimeoutError:
154+
log.warning("Disconnect timed out")
155+
except Exception as e:
156+
log.info("Close exception: %s %s", type(e).__name__, e)
157+
finally:
158+
self.writer = None
159+
160+
async def send(self, data: bytes, wait_for_marker: bool = True) -> None:
144161
async with self.lock:
145162
if self.writer is None or self.writer.is_closing():
146163
return None
147164

165+
if wait_for_marker:
166+
try:
167+
await asyncio.wait_for(self.marker_event.wait(), timeout=TCP.TIMEOUT)
168+
except asyncio.TimeoutError:
169+
raise TimeoutError
170+
148171
try:
149172
self.writer.write(data)
150173
await self.writer.drain()
@@ -162,7 +185,7 @@ async def recv(self, length: int = 0) -> Optional[bytes]:
162185
try:
163186
chunk = await asyncio.wait_for(
164187
self.reader.read(length - len(data)),
165-
TCP.TIMEOUT
188+
timeout=TCP.TIMEOUT
166189
)
167190
except (OSError, asyncio.TimeoutError):
168191
return None

pyrogram/connection/transport/tcp/tcp_abridged.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -26,12 +26,20 @@
2626

2727

2828
class TCPAbridged(TCP):
29-
def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEventLoop] = None) -> None:
30-
super().__init__(ipv6, proxy, loop)
29+
def __init__(
30+
self,
31+
ipv6: bool = False,
32+
proxy: Proxy = None,
33+
crypto_executor_workers: int = 1,
34+
loop: Optional[asyncio.AbstractEventLoop] = None,
35+
) -> None:
36+
super().__init__(ipv6, proxy, crypto_executor_workers, loop)
3137

3238
async def connect(self, address: Tuple[str, int]) -> None:
39+
self.marker_event.clear()
3340
await super().connect(address)
34-
await super().send(b"\xef")
41+
await super().send(b"\xef", wait_for_marker=False)
42+
self.marker_event.set()
3543

3644
async def send(self, data: bytes, *args) -> None:
3745
length = len(data) // 4

pyrogram/connection/transport/tcp/tcp_abridged_o.py

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -31,13 +31,20 @@
3131
class TCPAbridgedO(TCP):
3232
RESERVED = (b"HEAD", b"POST", b"GET ", b"OPTI", b"\xee" * 4)
3333

34-
def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEventLoop] = None) -> None:
35-
super().__init__(ipv6, proxy, loop)
34+
def __init__(
35+
self,
36+
ipv6: bool = False,
37+
proxy: Proxy = None,
38+
crypto_executor_workers: int = 1,
39+
loop: Optional[asyncio.AbstractEventLoop] = None,
40+
) -> None:
41+
super().__init__(ipv6, proxy, crypto_executor_workers, loop)
3642

3743
self.encrypt = None
3844
self.decrypt = None
3945

4046
async def connect(self, address: Tuple[str, int]) -> None:
47+
self.marker_event.clear()
4148
await super().connect(address)
4249

4350
while True:
@@ -54,12 +61,13 @@ async def connect(self, address: Tuple[str, int]) -> None:
5461

5562
nonce[56:64] = aes.ctr256_encrypt(nonce, *self.encrypt)[56:64]
5663

57-
await super().send(nonce)
64+
await super().send(nonce, wait_for_marker=False)
65+
self.marker_event.set()
5866

5967
async def send(self, data: bytes, *args) -> None:
6068
length = len(data) // 4
6169
data = (bytes([length]) if length <= 126 else b"\x7f" + length.to_bytes(3, "little")) + data
62-
payload = await self.loop.run_in_executor(pyrogram.crypto_executor, aes.ctr256_encrypt, data, *self.encrypt)
70+
payload = await self.loop.run_in_executor(self.crypto_executor, aes.ctr256_encrypt, data, *self.encrypt)
6371

6472
await super().send(payload)
6573

@@ -84,4 +92,4 @@ async def recv(self, length: int = 0) -> Optional[bytes]:
8492
if data is None:
8593
return None
8694

87-
return await self.loop.run_in_executor(pyrogram.crypto_executor, aes.ctr256_decrypt, data, *self.decrypt)
95+
return await self.loop.run_in_executor(self.crypto_executor, aes.ctr256_decrypt, data, *self.decrypt)

pyrogram/connection/transport/tcp/tcp_full.py

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -28,21 +28,29 @@
2828

2929

3030
class TCPFull(TCP):
31-
def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEventLoop] = None) -> None:
32-
super().__init__(ipv6, proxy, loop)
31+
def __init__(
32+
self,
33+
ipv6: bool = False,
34+
proxy: Proxy = None,
35+
crypto_executor_workers: int = 1,
36+
loop: Optional[asyncio.AbstractEventLoop] = None,
37+
) -> None:
38+
super().__init__(ipv6, proxy, crypto_executor_workers, loop)
3339

34-
self.seq_no: Optional[int] = None
40+
self.seq_no: int = 0
3541

3642
async def connect(self, address: Tuple[str, int]) -> None:
3743
await super().connect(address)
3844
self.seq_no = 0
3945

4046
async def send(self, data: bytes, *args) -> None:
47+
self.marker_event.clear()
4148
data = pack("<II", len(data) + 12, self.seq_no) + data
4249
data += pack("<I", crc32(data))
4350
self.seq_no += 1
4451

45-
await super().send(data)
52+
await super().send(data, wait_for_marker=False)
53+
self.marker_event.set()
4654

4755
async def recv(self, length: int = 0) -> Optional[bytes]:
4856
length = await super().recv(4)

pyrogram/connection/transport/tcp/tcp_intermediate.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,8 +31,10 @@ def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEven
3131
super().__init__(ipv6, proxy, loop)
3232

3333
async def connect(self, address: Tuple[str, int]) -> None:
34+
self.marker_event.clear()
3435
await super().connect(address)
35-
await super().send(b"\xee" * 4)
36+
await super().send(b"\xee" * 4, wait_for_marker=False)
37+
self.marker_event.set()
3638

3739
async def send(self, data: bytes, *args) -> None:
3840
await super().send(pack("<i", len(data)) + data)

pyrogram/connection/transport/tcp/tcp_intermediate_o.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ def __init__(self, ipv6: bool, proxy: Proxy, loop: Optional[asyncio.AbstractEven
3838
self.decrypt = None
3939

4040
async def connect(self, address: Tuple[str, int]) -> None:
41+
self.marker_event.clear()
4142
await super().connect(address)
4243

4344
while True:
@@ -54,7 +55,8 @@ async def connect(self, address: Tuple[str, int]) -> None:
5455

5556
nonce[56:64] = aes.ctr256_encrypt(nonce, *self.encrypt)[56:64]
5657

57-
await super().send(nonce)
58+
await super().send(nonce, wait_for_marker=False)
59+
self.marker_event.set()
5860

5961
async def send(self, data: bytes, *args) -> None:
6062
await super().send(

pyrogram/session/internals/msg_factory.py

Lines changed: 40 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -16,23 +16,49 @@
1616
# You should have received a copy of the GNU Lesser General Public License
1717
# along with Pyrogram. If not, see <http://www.gnu.org/licenses/>.
1818

19+
import asyncio
20+
import time
21+
1922
from pyrogram.raw.core import Message, MsgContainer, TLObject
2023
from pyrogram.raw.functions import Ping
21-
from pyrogram.raw.types import MsgsAck, HttpWait
22-
from .msg_id import MsgId
23-
from .seq_no import SeqNo
24-
25-
not_content_related = (Ping, HttpWait, MsgsAck, MsgContainer)
24+
from pyrogram.raw.types import HttpWait, MsgsAck
2625

2726

2827
class MsgFactory:
2928
def __init__(self):
30-
self.seq_no = SeqNo()
31-
32-
def __call__(self, body: TLObject) -> Message:
33-
return Message(
34-
body,
35-
MsgId(),
36-
self.seq_no(not isinstance(body, not_content_related)),
37-
len(body)
38-
)
29+
self._last_msg_id = 0
30+
31+
self._msg_id_lock = asyncio.Lock()
32+
self._seq_no_lock = asyncio.Lock()
33+
34+
self._content_related_messages_sent = 0
35+
36+
async def allocate_message_identity(self) -> int:
37+
async with self._msg_id_lock:
38+
now = time.time()
39+
40+
base_msg_id = int(now * (2**32)) & ~0b11
41+
42+
if base_msg_id <= self._last_msg_id:
43+
base_msg_id = self._last_msg_id + 4
44+
45+
self._last_msg_id = base_msg_id
46+
47+
return base_msg_id
48+
49+
async def allocate_message_sequence(self, is_content_related: bool) -> int:
50+
async with self._seq_no_lock:
51+
seq_no = (self._content_related_messages_sent * 2) + (1 if is_content_related else 0)
52+
53+
if is_content_related:
54+
self._content_related_messages_sent += 1
55+
56+
return seq_no
57+
58+
async def create(self, body: TLObject) -> Message:
59+
msg_id = await self.allocate_message_identity()
60+
61+
is_content_related = not isinstance(body, (Ping, HttpWait, MsgsAck, MsgContainer))
62+
seq_no = await self.allocate_message_sequence(is_content_related)
63+
64+
return Message(body, msg_id, seq_no, len(body))

pyrogram/session/internals/msg_id.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -23,13 +23,15 @@
2323

2424

2525
class MsgId:
26-
last_time = 0
27-
offset = 0
26+
_last_msg_id = 0
2827

2928
def __new__(cls) -> int:
30-
now = int(time.time())
31-
cls.offset = (cls.offset + 4) if now == cls.last_time else 0
32-
msg_id = (now * 2 ** 32) + cls.offset
33-
cls.last_time = now
29+
now = time.time()
30+
base_msg_id = int(now * (2**32)) & ~0b11
3431

35-
return msg_id
32+
if base_msg_id <= cls._last_msg_id:
33+
base_msg_id = cls._last_msg_id + 4
34+
35+
cls._last_msg_id = base_msg_id
36+
37+
return base_msg_id

0 commit comments

Comments
 (0)