Skip to content

Commit e49942d

Browse files
Use server address from session instead of hardcoded ips
1 parent c0ba58b commit e49942d

18 files changed

Lines changed: 382 additions & 247 deletions

pyrogram/client.py

Lines changed: 218 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -381,10 +381,8 @@ def __init__(
381381
self.business_connections = {}
382382

383383
self.sessions = {}
384-
self.sessions_lock = asyncio.Lock()
385-
386384
self.media_sessions = {}
387-
self.media_sessions_lock = asyncio.Lock()
385+
self.sessions_lock = asyncio.Lock()
388386

389387
self.save_file_semaphore = asyncio.Semaphore(self.max_concurrent_transmissions)
390388
self.get_file_semaphore = asyncio.Semaphore(self.max_concurrent_transmissions)
@@ -416,6 +414,8 @@ def __init__(
416414
else:
417415
self.loop = asyncio.get_event_loop()
418416

417+
self.__config = None
418+
419419
def __enter__(self):
420420
return self.start()
421421

@@ -898,12 +898,17 @@ async def load_session(self):
898898
await self.storage.api_id(self.api_id)
899899

900900
await self.storage.dc_id(2)
901+
await self.storage.server_address("149.154.167.51")
902+
await self.storage.port(443)
901903
await self.storage.date(0)
902904

903905
await self.storage.test_mode(self.test_mode)
904906
await self.storage.auth_key(
905907
await Auth(
906-
self, await self.storage.dc_id(),
908+
self,
909+
await self.storage.dc_id(),
910+
await self.storage.server_address(),
911+
await self.storage.port(),
907912
await self.storage.test_mode()
908913
).create()
909914
)
@@ -1136,10 +1141,20 @@ async def get_file(
11361141
try:
11371142
session = self.media_sessions.get(dc_id)
11381143
if not session:
1144+
dc_option = await self.get_dc_option(dc_id, is_media=True, ipv6=self.ipv6)
1145+
11391146
session = self.media_sessions[dc_id] = Session(
1140-
self, dc_id,
1141-
await Auth(self, dc_id, await self.storage.test_mode()).create()
1142-
if dc_id != await self.storage.dc_id()
1147+
self,
1148+
dc_id,
1149+
dc_option.ip_address,
1150+
dc_option.port,
1151+
await Auth(
1152+
self,
1153+
dc_id,
1154+
dc_option.ip_address,
1155+
dc_option.port,
1156+
await self.storage.test_mode()
1157+
).create() if dc_id != await self.storage.dc_id()
11431158
else await self.storage.auth_key(),
11441159
await self.storage.test_mode(),
11451160
is_media=True
@@ -1214,9 +1229,23 @@ async def get_file(
12141229
)
12151230

12161231
elif isinstance(r, raw.types.upload.FileCdnRedirect):
1232+
dc_option = await self.get_dc_option(dc_id, is_cdn=True, ipv6=self.ipv6)
1233+
12171234
cdn_session = Session(
1218-
self, r.dc_id, await Auth(self, r.dc_id, await self.storage.test_mode()).create(),
1219-
await self.storage.test_mode(), is_media=True, is_cdn=True
1235+
self,
1236+
r.dc_id,
1237+
dc_option.ip_address,
1238+
dc_option.port,
1239+
await Auth(
1240+
self,
1241+
r.dc_id,
1242+
dc_option.ip_address,
1243+
dc_option.port,
1244+
await self.storage.test_mode()
1245+
).create(),
1246+
await self.storage.test_mode(),
1247+
is_media=True,
1248+
is_cdn=True
12201249
)
12211250

12221251
try:
@@ -1302,6 +1331,186 @@ async def get_file(
13021331
except Exception as e:
13031332
log.exception(e)
13041333

1334+
async def get_session(
1335+
self,
1336+
dc_id: int,
1337+
is_media: Optional[bool] = False,
1338+
export_authorization: Optional[bool] = True,
1339+
server_address: Optional[str] = None,
1340+
port: Optional[int] = None
1341+
) -> "Session":
1342+
"""Get existing session or create a new one.
1343+
1344+
Parameters:
1345+
dc_id (``int``):
1346+
Datacenter identifier.
1347+
1348+
is_media (``bool``, *optional*):
1349+
Pass True to get or create a media session.
1350+
1351+
export_authorization (``bool``, *optional*):
1352+
Pass True to export authorization after creating the session.
1353+
Used only when creating a new session.
1354+
1355+
server_address (``str``, *optional*):
1356+
Custom server address to connect to.
1357+
Used only when creating a new session.
1358+
1359+
port (``int``, *optional*):
1360+
Custom port to connect to.
1361+
Used only when creating a new session.
1362+
"""
1363+
if dc_id == await self.storage.dc_id():
1364+
return self.session
1365+
1366+
sessions = self.media_sessions if is_media else self.sessions
1367+
1368+
async with self.sessions_lock:
1369+
if sessions.get(dc_id):
1370+
return sessions[dc_id]
1371+
1372+
dc_option = await self.get_dc_option(dc_id, is_media=is_media, ipv6=self.ipv6)
1373+
1374+
session = self.media_sessions[dc_id] = Session(
1375+
self,
1376+
dc_id,
1377+
server_address or dc_option.ip_address,
1378+
port or dc_option.port,
1379+
await Auth(
1380+
self,
1381+
dc_id,
1382+
server_address or dc_option.ip_address,
1383+
port or dc_option.port,
1384+
await self.storage.test_mode()
1385+
).create(),
1386+
await self.storage.test_mode(), is_media=is_media
1387+
)
1388+
1389+
await session.start()
1390+
1391+
if export_authorization:
1392+
for _ in range(3):
1393+
exported_auth = await self.invoke(
1394+
raw.functions.auth.ExportAuthorization(
1395+
dc_id=dc_id
1396+
)
1397+
)
1398+
1399+
try:
1400+
await session.invoke(
1401+
raw.functions.auth.ImportAuthorization(
1402+
id=exported_auth.id,
1403+
bytes=exported_auth.bytes
1404+
)
1405+
)
1406+
except AuthBytesInvalid:
1407+
continue
1408+
else:
1409+
break
1410+
else:
1411+
await session.stop()
1412+
raise AuthBytesInvalid
1413+
1414+
return session
1415+
1416+
async def get_dc_option(
1417+
self,
1418+
dc_id: int,
1419+
is_media: bool = False,
1420+
is_cdn: bool = False,
1421+
ipv6: bool = False
1422+
) -> "raw.types.DcOption":
1423+
if not self.__config:
1424+
self.__config = await self.invoke(raw.functions.help.GetConfig())
1425+
1426+
options = [dc for dc in self.__config.dc_options if dc.id == dc_id and dc.ipv6 == ipv6] # type: List[raw.types.DcOption]
1427+
1428+
if not options:
1429+
raise ValueError(f"DC{dc_id} not found")
1430+
1431+
if is_cdn:
1432+
cdn_options = [dc for dc in options if dc.cdn]
1433+
1434+
if cdn_options:
1435+
return cdn_options[0]
1436+
1437+
log.debug(
1438+
"No CDN datacenter found for DC%s, falling back to prod DC",
1439+
dc_id
1440+
)
1441+
1442+
is_media = True
1443+
1444+
if is_media:
1445+
media_options = [dc for dc in options if dc.media_only]
1446+
1447+
if media_options:
1448+
return media_options[0]
1449+
1450+
log.debug(
1451+
"No media datacenter found for DC%s, falling back to prod DC",
1452+
dc_id
1453+
)
1454+
1455+
prod_options = [dc for dc in options if not dc.media_only]
1456+
1457+
if prod_options:
1458+
return prod_options[0]
1459+
1460+
raise ValueError("No suitable DC found")
1461+
1462+
async def set_dc(
1463+
self,
1464+
dc_id: Optional[int] = None,
1465+
*,
1466+
server_address: Optional[str] = None,
1467+
port: Optional[int] = None
1468+
):
1469+
"""Be careful with this method, you can easily break your session."""
1470+
if not self.__config:
1471+
self.__config = await self.invoke(raw.functions.help.GetConfig())
1472+
1473+
dc_id = dc_id or self.__config.this_dc
1474+
1475+
dc_option = await self.get_dc_option(dc_id, ipv6=self.ipv6)
1476+
1477+
server_address = server_address or dc_option.ip_address
1478+
port = port or dc_option.port
1479+
1480+
if dc_id == self.__config.this_dc and (self.session.server_address != server_address or self.session.port != port):
1481+
await self.storage.server_address(server_address)
1482+
await self.storage.port(port)
1483+
1484+
self.session.server_address = await self.storage.server_address()
1485+
self.session.port = await self.storage.port()
1486+
1487+
await self.session.restart()
1488+
log.info("Changed the current session DC%s address to %s:%s", dc_id, server_address, port)
1489+
else:
1490+
prod_session = self.sessions.get(dc_id)
1491+
1492+
if prod_session and (prod_session.server_address != server_address or prod_session.port != port):
1493+
prod_session.server_address = server_address
1494+
prod_session.port = port
1495+
1496+
await prod_session.restart()
1497+
log.info("Changed session DC%s address to %s:%s", dc_id, server_address, port)
1498+
else:
1499+
await self.get_session(dc_id, server_address=server_address, port=port)
1500+
log.info("Created new session DC%s with address %s:%s", dc_id, server_address, port)
1501+
1502+
media_session = self.media_sessions.get(dc_id)
1503+
1504+
if media_session and (media_session.server_address != server_address or media_session.port != port):
1505+
media_session.server_address = server_address
1506+
media_session.port = port
1507+
1508+
await media_session.restart()
1509+
log.info("Changed session DC%s (media) address to %s:%s", dc_id, server_address, port)
1510+
else:
1511+
await self.get_session(dc_id, is_media=True, server_address=server_address, port=port)
1512+
log.info("Created new session DC%s (media) with address %s:%s", dc_id, server_address, port)
1513+
13051514
def guess_mime_type(self, filename: Union[str, BytesIO]) -> Optional[str]:
13061515
if isinstance(filename, BytesIO):
13071516
return self.mimetypes.guess_type(filename.name)[0]

pyrogram/connection/connection.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121
from typing import Optional, Type
2222

2323
from .transport import TCP, TCPAbridged
24-
from ..session.internals import DataCenter
2524

2625
log = logging.getLogger(__name__)
2726

@@ -32,6 +31,8 @@ class Connection:
3231
def __init__(
3332
self,
3433
dc_id: int,
34+
server_address: str,
35+
port: int,
3536
test_mode: bool,
3637
ipv6: bool,
3738
proxy: dict,
@@ -41,14 +42,15 @@ def __init__(
4142
loop: Optional[asyncio.AbstractEventLoop] = None
4243
) -> None:
4344
self.dc_id = dc_id
45+
self.server_address = server_address
46+
self.port = port
4447
self.test_mode = test_mode
4548
self.ipv6 = ipv6
4649
self.proxy = proxy
4750
self.media = media
4851
self.protocol_factory = protocol_factory
4952
self.crypto_executor_workers = crypto_executor_workers
5053

51-
self.address = DataCenter(dc_id, test_mode, ipv6, media)
5254
self.protocol: Optional[TCP] = None
5355

5456
if isinstance(loop, asyncio.AbstractEventLoop):
@@ -62,7 +64,7 @@ async def connect(self) -> None:
6264

6365
try:
6466
log.info("Connecting...")
65-
await self.protocol.connect(self.address)
67+
await self.protocol.connect((self.server_address, self.port))
6668
except OSError as e:
6769
log.warning("Unable to connect due to network issues: %s", e)
6870
await self.protocol.close()

pyrogram/methods/advanced/save_file.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -145,9 +145,16 @@ async def worker(session):
145145
dc_id = await self.storage.dc_id()
146146

147147
session = self.media_sessions.get(dc_id)
148+
148149
if not session:
150+
dc_option = await self.get_dc_option(dc_id, is_media=True, ipv6=self.ipv6)
151+
149152
session = self.media_sessions[dc_id] = Session(
150-
self, dc_id, await self.storage.auth_key(),
153+
self,
154+
dc_id,
155+
dc_option.ip_address,
156+
dc_option.port,
157+
await self.storage.auth_key(),
151158
await self.storage.test_mode(), is_media=True
152159
)
153160
await session.start()

pyrogram/methods/auth/connect.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -40,8 +40,12 @@ async def connect(
4040
await self.load_session()
4141

4242
self.session = Session(
43-
self, await self.storage.dc_id(),
44-
await self.storage.auth_key(), await self.storage.test_mode()
43+
self,
44+
await self.storage.dc_id(),
45+
await self.storage.server_address(),
46+
await self.storage.port(),
47+
await self.storage.auth_key(),
48+
await self.storage.test_mode()
4549
)
4650

4751
await self.session.start()

pyrogram/methods/auth/send_code.py

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
from pyrogram import raw
2424
from pyrogram import types
2525
from pyrogram.errors import PhoneMigrate, NetworkMigrate
26+
from pyrogram.raw.base import dc_option
2627
from pyrogram.session import Session, Auth
2728

2829
log = logging.getLogger(__name__)
@@ -102,18 +103,30 @@ async def send_code(
102103
)
103104
)
104105
except (PhoneMigrate, NetworkMigrate) as e:
106+
dc_option = await self.get_dc_option(e.value, ipv6=self.ipv6)
105107
await self.session.stop()
106108

107109
await self.storage.dc_id(e.value)
110+
await self.storage.server_address(dc_option.ip_address)
111+
await self.storage.port(dc_option.port)
112+
108113
await self.storage.auth_key(
109114
await Auth(
110-
self, await self.storage.dc_id(),
115+
self,
116+
await self.storage.dc_id(),
117+
await self.storage.server_address(),
118+
await self.storage.port(),
111119
await self.storage.test_mode()
112120
).create()
113121
)
122+
114123
self.session = Session(
115-
self, await self.storage.dc_id(),
116-
await self.storage.auth_key(), await self.storage.test_mode()
124+
self,
125+
await self.storage.dc_id(),
126+
await self.storage.server_address(),
127+
await self.storage.port(),
128+
await self.storage.auth_key(),
129+
await self.storage.test_mode()
117130
)
118131

119132
await self.session.start()

0 commit comments

Comments
 (0)