maxconn 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.
maxconn/__init__.py ADDED
@@ -0,0 +1,68 @@
1
+ """maxconn - SSH and Telnet client built from scratch."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from maxconn.exceptions import (
6
+ AuthenticationError,
7
+ ChannelError,
8
+ ConnectionTimeoutError,
9
+ MaxConnError,
10
+ ProtocolError,
11
+ )
12
+ from maxconn.transport.base import CommandResult, Connection, Transport
13
+
14
+ __all__ = [
15
+ "AuthenticationError",
16
+ "ChannelError",
17
+ "CommandResult",
18
+ "Connection",
19
+ "ConnectionTimeoutError",
20
+ "MaxConnError",
21
+ "ProtocolError",
22
+ "connect",
23
+ ]
24
+
25
+ _DEFAULT_PORTS = {"telnet": 23, "ssh": 22}
26
+
27
+
28
+ def connect(
29
+ host: str,
30
+ *,
31
+ protocol: str,
32
+ username: str,
33
+ password: str | None = None,
34
+ pkey: object | None = None,
35
+ port: int | None = None,
36
+ timeout: float = 10.0,
37
+ connect_timeout: float | None = None,
38
+ auth_timeout: float | None = None,
39
+ command_timeout: float = 5.0,
40
+ prompt_timeout: float = 10.0,
41
+ ) -> Connection:
42
+ # Imported per-protocol, not at module load, so using one transport
43
+ # never pulls in the other's dependencies (see the "import only what
44
+ # you use" rule in the project brief).
45
+ transport: Transport
46
+ if protocol == "telnet":
47
+ from maxconn.transport.telnet.transport import TelnetTransport
48
+
49
+ transport = TelnetTransport()
50
+ elif protocol == "ssh":
51
+ from maxconn.transport.ssh.transport import SSHTransport
52
+
53
+ transport = SSHTransport()
54
+ else:
55
+ raise ValueError(f"Unsupported protocol: {protocol!r}. Supported: 'telnet', 'ssh'")
56
+
57
+ resolved_port = port if port is not None else _DEFAULT_PORTS[protocol]
58
+ resolved_connect_timeout = connect_timeout if connect_timeout is not None else timeout
59
+ resolved_auth_timeout = auth_timeout if auth_timeout is not None else timeout
60
+ transport.connect(host, resolved_port, resolved_connect_timeout)
61
+ transport.authenticate(username, password=password, pkey=pkey, timeout=resolved_auth_timeout)
62
+ return Connection(
63
+ transport,
64
+ host=host,
65
+ protocol=protocol,
66
+ command_timeout=command_timeout,
67
+ prompt_timeout=prompt_timeout,
68
+ )
@@ -0,0 +1,5 @@
1
+ """Prompt/expect helpers for interactive network CLIs."""
2
+
3
+ from maxconn.automation.expect import ExpectSession, PromptProfile
4
+
5
+ __all__ = ["ExpectSession", "PromptProfile"]
@@ -0,0 +1,84 @@
1
+ """Small expect-style command runner for prompt-based CLIs."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import time
6
+ from enum import Enum
7
+ from typing import Protocol
8
+
9
+ from maxconn.exceptions import ConnectionTimeoutError
10
+
11
+
12
+ class PromptProfile(Enum):
13
+ GENERIC = ("generic", (">", "#"))
14
+ CISCO = ("cisco", (">", "#"))
15
+ HUAWEI = ("huawei", ("<", ">", "]"))
16
+
17
+ @property
18
+ def markers(self) -> tuple[str, ...]:
19
+ return self.value[1]
20
+
21
+
22
+ class ExpectConnection(Protocol):
23
+ def send(self, data: bytes | str) -> None: ...
24
+
25
+ def recv(self, timeout: float | None = None) -> bytes: ...
26
+
27
+
28
+ class ExpectSession:
29
+ def __init__(
30
+ self,
31
+ connection: ExpectConnection,
32
+ *,
33
+ prompt_markers: tuple[str, ...] | PromptProfile,
34
+ pagination_markers: tuple[str, ...] = (),
35
+ ) -> None:
36
+ self._connection = connection
37
+ self._prompt_markers = (
38
+ prompt_markers.markers if isinstance(prompt_markers, PromptProfile) else prompt_markers
39
+ )
40
+ self._pagination_markers = pagination_markers
41
+
42
+ def run(self, command: str, *, timeout: float = 10.0, strip_echo: bool = False) -> str:
43
+ self._connection.send(command + "\n")
44
+ output = self._read_until_prompt(timeout)
45
+ if strip_echo:
46
+ output = self._strip_command_echo(output, command)
47
+ return output
48
+
49
+ def _read_until_prompt(self, timeout: float) -> str:
50
+ deadline = time.monotonic() + timeout
51
+ buffer = ""
52
+ while time.monotonic() < deadline:
53
+ remaining = deadline - time.monotonic()
54
+ try:
55
+ chunk = self._connection.recv(timeout=max(remaining, 0.01))
56
+ except ConnectionTimeoutError as exc:
57
+ raise ConnectionTimeoutError(
58
+ f"Timed out waiting for prompt {self._prompt_markers!r}; got: {buffer!r}"
59
+ ) from exc
60
+
61
+ buffer += chunk.decode(errors="replace")
62
+ buffer = self._answer_pagination(buffer)
63
+ if any(marker in buffer for marker in self._prompt_markers):
64
+ return buffer
65
+
66
+ raise ConnectionTimeoutError(
67
+ f"Timed out waiting for prompt {self._prompt_markers!r}; got: {buffer!r}"
68
+ )
69
+
70
+ def _answer_pagination(self, buffer: str) -> str:
71
+ for marker in self._pagination_markers:
72
+ if marker in buffer:
73
+ self._connection.send(" ")
74
+ buffer = buffer.replace(marker, "")
75
+ return buffer
76
+
77
+ @staticmethod
78
+ def _strip_command_echo(output: str, command: str) -> str:
79
+ normalized_command = command.strip()
80
+ for separator in ("\r\n", "\n"):
81
+ prefix = normalized_command + separator
82
+ if output.startswith(prefix):
83
+ return output[len(prefix) :]
84
+ return output
maxconn/exceptions.py ADDED
@@ -0,0 +1,18 @@
1
+ class MaxConnError(Exception):
2
+ """Base exception for all maxconn errors."""
3
+
4
+
5
+ class ConnectionTimeoutError(MaxConnError):
6
+ """Raised when a connection attempt or a read times out."""
7
+
8
+
9
+ class AuthenticationError(MaxConnError):
10
+ """Raised when authentication with the remote host fails."""
11
+
12
+
13
+ class ProtocolError(MaxConnError):
14
+ """Raised when a malformed or unexpected protocol message is received."""
15
+
16
+
17
+ class ChannelError(MaxConnError):
18
+ """Raised when a channel operation fails (SSH channels)."""
File without changes
@@ -0,0 +1,162 @@
1
+ """Transport abstraction shared by SSH and Telnet implementations."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import logging
6
+ import re
7
+ import time
8
+ from abc import ABC, abstractmethod
9
+ from dataclasses import dataclass
10
+
11
+ from maxconn.automation import ExpectSession, PromptProfile
12
+ from maxconn.exceptions import ConnectionTimeoutError
13
+
14
+ _AUDIT_LOG = logging.getLogger("maxconn.audit")
15
+ _SECRET_PATTERNS = (
16
+ re.compile(r"(?i)\b(password|passwd|secret|token|key)\s+\S+"),
17
+ re.compile(r"(?i)(--password|--passwd|--secret|--token|--key)\s+\S+"),
18
+ )
19
+
20
+
21
+ class Transport(ABC):
22
+ @abstractmethod
23
+ def connect(self, host: str, port: int, timeout: float) -> None: ...
24
+
25
+ @abstractmethod
26
+ def authenticate(
27
+ self,
28
+ username: str,
29
+ password: str | None = None,
30
+ pkey: object | None = None,
31
+ timeout: float | None = None,
32
+ ) -> None: ...
33
+ # `pkey`'s concrete type is transport-specific (e.g. SSHTransport
34
+ # expects a cryptography RSAPrivateKey); transports that don't support
35
+ # key-based auth raise AuthenticationError when it's passed.
36
+
37
+ @abstractmethod
38
+ def send(self, data: bytes | str) -> None: ...
39
+
40
+ @abstractmethod
41
+ def recv(self, timeout: float | None = None) -> bytes: ...
42
+
43
+ @abstractmethod
44
+ def close(self) -> None: ...
45
+
46
+
47
+ def read_until(transport: Transport, markers: tuple[str, ...], timeout: float) -> str:
48
+ deadline = time.monotonic() + timeout
49
+ buffer = ""
50
+ while time.monotonic() < deadline:
51
+ remaining = deadline - time.monotonic()
52
+ chunk = transport.recv(timeout=max(remaining, 0.01))
53
+ buffer += chunk.decode(errors="replace")
54
+ if any(marker.lower() in buffer.lower() for marker in markers):
55
+ return buffer
56
+ raise ConnectionTimeoutError(f"Timed out waiting for markers {markers!r}; got: {buffer!r}")
57
+
58
+
59
+ @dataclass(frozen=True)
60
+ class CommandResult:
61
+ command: str
62
+ text: str
63
+ bytes: bytes
64
+ elapsed: float
65
+ exit_status: int | None = None
66
+
67
+ @property
68
+ def ok(self) -> bool:
69
+ return self.exit_status in (None, 0)
70
+
71
+
72
+ class Connection:
73
+ def __init__(
74
+ self,
75
+ transport: Transport,
76
+ *,
77
+ host: str | None = None,
78
+ protocol: str | None = None,
79
+ command_timeout: float = 5.0,
80
+ prompt_timeout: float = 10.0,
81
+ prompt_profile: PromptProfile = PromptProfile.GENERIC,
82
+ ) -> None:
83
+ self._transport = transport
84
+ self.host = host
85
+ self.protocol = protocol
86
+ self.command_timeout = command_timeout
87
+ self.prompt_timeout = prompt_timeout
88
+ self.prompt_profile = prompt_profile
89
+
90
+ def send_command(self, command: str, *, read_timeout: float = 5.0) -> str:
91
+ self._transport.send(command + "\n")
92
+ deadline = time.monotonic() + read_timeout
93
+ buffer = ""
94
+ while time.monotonic() < deadline:
95
+ remaining = deadline - time.monotonic()
96
+ try:
97
+ chunk = self._transport.recv(timeout=max(remaining, 0.01))
98
+ except ConnectionTimeoutError:
99
+ break
100
+ buffer += chunk.decode(errors="replace")
101
+ return buffer
102
+
103
+ def run(
104
+ self,
105
+ command: str,
106
+ *,
107
+ prompt_markers: tuple[str, ...] | None = None,
108
+ timeout: float | None = None,
109
+ strip_echo: bool = True,
110
+ pagination_markers: tuple[str, ...] = ("--More--",),
111
+ ) -> CommandResult:
112
+ start = time.monotonic()
113
+ resolved_timeout = timeout if timeout is not None else self.command_timeout
114
+ resolved_markers = prompt_markers if prompt_markers is not None else self.prompt_profile.markers
115
+ expect = ExpectSession(
116
+ self,
117
+ prompt_markers=resolved_markers,
118
+ pagination_markers=pagination_markers,
119
+ )
120
+
121
+ text = expect.run(command, timeout=resolved_timeout, strip_echo=strip_echo)
122
+ elapsed = time.monotonic() - start
123
+ result = CommandResult(
124
+ command=command,
125
+ text=text,
126
+ bytes=text.encode(),
127
+ elapsed=elapsed,
128
+ )
129
+ _AUDIT_LOG.info(
130
+ "command completed host=%s protocol=%s command=%r elapsed=%.3f ok=%s",
131
+ self.host,
132
+ self.protocol,
133
+ _redact(command),
134
+ elapsed,
135
+ result.ok,
136
+ )
137
+ return result
138
+
139
+ def send(self, data: bytes | str) -> None:
140
+ self._transport.send(data)
141
+
142
+ def recv(self, timeout: float | None = None) -> bytes:
143
+ return self._transport.recv(timeout=timeout)
144
+
145
+ def read_until(self, marker: str, timeout: float = 10.0) -> str:
146
+ return read_until(self._transport, (marker,), timeout)
147
+
148
+ def close(self) -> None:
149
+ self._transport.close()
150
+
151
+ def __enter__(self) -> Connection:
152
+ return self
153
+
154
+ def __exit__(self, exc_type, exc, tb) -> None:
155
+ self.close()
156
+
157
+
158
+ def _redact(value: str) -> str:
159
+ redacted = value
160
+ for pattern in _SECRET_PATTERNS:
161
+ redacted = pattern.sub(lambda match: f"{match.group(1)} <redacted>", redacted)
162
+ return redacted
File without changes
@@ -0,0 +1,104 @@
1
+ """SSH user authentication (RFC 4252): the ssh-userauth service request,
2
+ plus the "password" and "publickey" authentication methods."""
3
+
4
+ from __future__ import annotations
5
+
6
+ from cryptography.hazmat.primitives import hashes
7
+ from cryptography.hazmat.primitives.asymmetric import padding, rsa
8
+
9
+ from maxconn.exceptions import AuthenticationError, ProtocolError
10
+ from maxconn.transport.ssh import messages
11
+ from maxconn.transport.ssh.negotiate import EncryptedSession
12
+ from maxconn.transport.ssh.wire import Reader, encode_mpint, encode_string
13
+
14
+ _MAX_BANNERS = 10
15
+ _PUBKEY_ALGORITHM = "rsa-sha2-256"
16
+
17
+
18
+ def build_rsa_public_key_blob(public_key: rsa.RSAPublicKey) -> bytes:
19
+ numbers = public_key.public_numbers()
20
+ return encode_string(b"ssh-rsa") + encode_mpint(numbers.e) + encode_mpint(numbers.n)
21
+
22
+
23
+ def request_userauth_service(session: EncryptedSession) -> None:
24
+ request = bytes([messages.SSH_MSG_SERVICE_REQUEST]) + encode_string(b"ssh-userauth")
25
+ session.send_message(request)
26
+
27
+ reader = Reader(session.recv_message())
28
+ msg_type = reader.read_byte()
29
+ if msg_type != messages.SSH_MSG_SERVICE_ACCEPT:
30
+ raise ProtocolError(f"Expected SSH_MSG_SERVICE_ACCEPT ({messages.SSH_MSG_SERVICE_ACCEPT}), got {msg_type}")
31
+ service_name = reader.read_string()
32
+ if service_name != b"ssh-userauth":
33
+ raise ProtocolError(f"Unexpected service accepted: {service_name!r}")
34
+
35
+
36
+ def _read_until_auth_result(session: EncryptedSession) -> bytes:
37
+ """SSH_MSG_USERAUTH_BANNER may arrive before the real result; skip it."""
38
+ for _ in range(_MAX_BANNERS):
39
+ payload = session.recv_message()
40
+ if payload and payload[0] == messages.SSH_MSG_USERAUTH_BANNER:
41
+ continue
42
+ return payload
43
+ raise ProtocolError("Too many SSH_MSG_USERAUTH_BANNER messages without an auth result")
44
+
45
+
46
+ def authenticate_password(session: EncryptedSession, username: str, password: str) -> None:
47
+ request = (
48
+ bytes([messages.SSH_MSG_USERAUTH_REQUEST])
49
+ + encode_string(username.encode("utf-8"))
50
+ + encode_string(b"ssh-connection")
51
+ + encode_string(b"password")
52
+ + bytes([0]) # FALSE: this is not a change-password request
53
+ + encode_string(password.encode("utf-8"))
54
+ )
55
+ session.send_message(request)
56
+
57
+ payload = _read_until_auth_result(session)
58
+ reader = Reader(payload)
59
+ msg_type = reader.read_byte()
60
+ if msg_type == messages.SSH_MSG_USERAUTH_SUCCESS:
61
+ return
62
+ if msg_type == messages.SSH_MSG_USERAUTH_FAILURE:
63
+ raise AuthenticationError(f"SSH password authentication failed for user {username!r}")
64
+ raise ProtocolError(f"Unexpected message during password authentication: {msg_type}")
65
+
66
+
67
+ def authenticate_publickey(session: EncryptedSession, username: str, private_key: rsa.RSAPrivateKey) -> None:
68
+ """RFC 4252 §7. Signs directly (no accept-probing round trip) since the
69
+ servers we target don't require it."""
70
+ public_key_blob = build_rsa_public_key_blob(private_key.public_key())
71
+
72
+ signed_part = (
73
+ encode_string(session.session_id)
74
+ + bytes([messages.SSH_MSG_USERAUTH_REQUEST])
75
+ + encode_string(username.encode("utf-8"))
76
+ + encode_string(b"ssh-connection")
77
+ + encode_string(b"publickey")
78
+ + bytes([1]) # TRUE: a signature is included
79
+ + encode_string(_PUBKEY_ALGORITHM.encode("ascii"))
80
+ + encode_string(public_key_blob)
81
+ )
82
+ signature = private_key.sign(signed_part, padding.PKCS1v15(), hashes.SHA256())
83
+ signature_blob = encode_string(_PUBKEY_ALGORITHM.encode("ascii")) + encode_string(signature)
84
+
85
+ request = (
86
+ bytes([messages.SSH_MSG_USERAUTH_REQUEST])
87
+ + encode_string(username.encode("utf-8"))
88
+ + encode_string(b"ssh-connection")
89
+ + encode_string(b"publickey")
90
+ + bytes([1])
91
+ + encode_string(_PUBKEY_ALGORITHM.encode("ascii"))
92
+ + encode_string(public_key_blob)
93
+ + encode_string(signature_blob)
94
+ )
95
+ session.send_message(request)
96
+
97
+ payload = _read_until_auth_result(session)
98
+ reader = Reader(payload)
99
+ msg_type = reader.read_byte()
100
+ if msg_type == messages.SSH_MSG_USERAUTH_SUCCESS:
101
+ return
102
+ if msg_type == messages.SSH_MSG_USERAUTH_FAILURE:
103
+ raise AuthenticationError(f"SSH public-key authentication failed for user {username!r}")
104
+ raise ProtocolError(f"Unexpected message during public-key authentication: {msg_type}")
@@ -0,0 +1,131 @@
1
+ """SSH channel handling: open a session channel and run exec/shell commands
2
+ (RFC 4254 sections 4-6). v0.1 does not track outgoing flow control (channel
3
+ window replenishment for data *we* send) - fine for interactive command/CLI
4
+ use, not for bulk transfers."""
5
+
6
+ from __future__ import annotations
7
+
8
+ from maxconn.exceptions import ChannelError
9
+ from maxconn.transport.ssh import messages
10
+ from maxconn.transport.ssh.negotiate import EncryptedSession
11
+ from maxconn.transport.ssh.wire import Reader, encode_string, encode_uint32
12
+
13
+ INITIAL_WINDOW_SIZE = 2 * 1024 * 1024
14
+ MAX_PACKET_SIZE = 32768
15
+
16
+ _CONTROL_ONLY_MESSAGES = (
17
+ messages.SSH_MSG_CHANNEL_WINDOW_ADJUST,
18
+ messages.SSH_MSG_CHANNEL_SUCCESS,
19
+ messages.SSH_MSG_CHANNEL_FAILURE,
20
+ messages.SSH_MSG_CHANNEL_REQUEST, # e.g. "exit-status" - informational only
21
+ messages.SSH_MSG_CHANNEL_EOF,
22
+ )
23
+
24
+
25
+ class SSHChannel:
26
+ def __init__(self, session: EncryptedSession, local_id: int, peer_id: int) -> None:
27
+ self._session = session
28
+ self.local_id = local_id
29
+ self.peer_id = peer_id
30
+ self._closed = False
31
+
32
+ def request_exec(self, command: str) -> None:
33
+ payload = (
34
+ bytes([messages.SSH_MSG_CHANNEL_REQUEST])
35
+ + encode_uint32(self.peer_id)
36
+ + encode_string(b"exec")
37
+ + bytes([0]) # want_reply = False: caller reads the command output instead
38
+ + encode_string(command.encode("utf-8"))
39
+ )
40
+ self._session.send_message(payload)
41
+
42
+ def request_shell(self) -> None:
43
+ pty_payload = (
44
+ bytes([messages.SSH_MSG_CHANNEL_REQUEST])
45
+ + encode_uint32(self.peer_id)
46
+ + encode_string(b"pty-req")
47
+ + bytes([0])
48
+ + encode_string(b"xterm")
49
+ + encode_uint32(80)
50
+ + encode_uint32(24)
51
+ + encode_uint32(640)
52
+ + encode_uint32(480)
53
+ + encode_string(b"")
54
+ )
55
+ self._session.send_message(pty_payload)
56
+
57
+ shell_payload = (
58
+ bytes([messages.SSH_MSG_CHANNEL_REQUEST])
59
+ + encode_uint32(self.peer_id)
60
+ + encode_string(b"shell")
61
+ + bytes([0])
62
+ )
63
+ self._session.send_message(shell_payload)
64
+
65
+ def send_data(self, data: bytes) -> None:
66
+ payload = bytes([messages.SSH_MSG_CHANNEL_DATA]) + encode_uint32(self.peer_id) + encode_string(data)
67
+ self._session.send_message(payload)
68
+
69
+ def recv_data(self) -> bytes:
70
+ """Read one channel message and return its data payload (empty
71
+ bytes for control-only messages like a window adjustment)."""
72
+ while True:
73
+ payload = self._session.recv_message()
74
+ reader = Reader(payload)
75
+ msg_type = reader.read_byte()
76
+
77
+ if msg_type == messages.SSH_MSG_CHANNEL_DATA:
78
+ reader.read_uint32() # recipient channel (ours; unused)
79
+ return reader.read_string()
80
+ if msg_type == messages.SSH_MSG_CHANNEL_EXTENDED_DATA:
81
+ reader.read_uint32()
82
+ reader.read_uint32() # data_type_code (e.g. stderr)
83
+ return reader.read_string()
84
+ if msg_type in _CONTROL_ONLY_MESSAGES:
85
+ continue
86
+ if msg_type == messages.SSH_MSG_CHANNEL_CLOSE:
87
+ self._closed = True
88
+ close_payload = bytes([messages.SSH_MSG_CHANNEL_CLOSE]) + encode_uint32(self.peer_id)
89
+ self._session.send_message(close_payload)
90
+ return b""
91
+
92
+ raise ChannelError(f"Unexpected message on SSH channel: type {msg_type}")
93
+
94
+ @property
95
+ def closed(self) -> bool:
96
+ return self._closed
97
+
98
+ def close(self) -> None:
99
+ if not self._closed:
100
+ self._session.send_message(bytes([messages.SSH_MSG_CHANNEL_CLOSE]) + encode_uint32(self.peer_id))
101
+ self._closed = True
102
+
103
+
104
+ def open_session_channel(session: EncryptedSession, local_id: int = 0) -> SSHChannel:
105
+ request = (
106
+ bytes([messages.SSH_MSG_CHANNEL_OPEN])
107
+ + encode_string(b"session")
108
+ + encode_uint32(local_id)
109
+ + encode_uint32(INITIAL_WINDOW_SIZE)
110
+ + encode_uint32(MAX_PACKET_SIZE)
111
+ )
112
+ session.send_message(request)
113
+
114
+ payload = session.recv_message()
115
+ reader = Reader(payload)
116
+ msg_type = reader.read_byte()
117
+
118
+ if msg_type == messages.SSH_MSG_CHANNEL_OPEN_FAILURE:
119
+ reader.read_uint32() # recipient channel
120
+ reason_code = reader.read_uint32()
121
+ description = reader.read_string().decode(errors="replace")
122
+ raise ChannelError(f"SSH channel open failed (reason {reason_code}): {description}")
123
+ if msg_type != messages.SSH_MSG_CHANNEL_OPEN_CONFIRMATION:
124
+ raise ChannelError(f"Unexpected reply to SSH_MSG_CHANNEL_OPEN: type {msg_type}")
125
+
126
+ reader.read_uint32() # recipient channel (== local_id, echoed back)
127
+ peer_id = reader.read_uint32()
128
+ reader.read_uint32() # peer's initial window size (unused, see module docstring)
129
+ reader.read_uint32() # peer's max packet size (unused, same reason)
130
+
131
+ return SSHChannel(session, local_id, peer_id)
@@ -0,0 +1,105 @@
1
+ """Diffie-Hellman group14-sha256 key exchange (RFC 4253 section 8, group
2
+ from RFC 3526 section 3 - the 2048-bit MODP "group 14")."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import hashlib
7
+ import secrets
8
+ import socket
9
+ from dataclasses import dataclass
10
+
11
+ from maxconn.exceptions import ProtocolError
12
+ from maxconn.transport.ssh import messages
13
+ from maxconn.transport.ssh.packet import decode_binary_packet, encode_binary_packet
14
+ from maxconn.transport.ssh.socket_reader import SocketReader
15
+ from maxconn.transport.ssh.wire import Reader, encode_mpint, encode_string
16
+
17
+ # RFC 3526 section 3, 2048-bit MODP Group ("group 14").
18
+ _P_HEX = (
19
+ "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD"
20
+ "129024E088A67CC74020BBEA63B139B22514A08798E3404"
21
+ "DDEF9519B3CD3A431B302B0A6DF25F14374FE1356D6D51C"
22
+ "245E485B576625E7EC6F44C42E9A637ED6B0BFF5CB6F406"
23
+ "B7EDEE386BFB5A899FA5AE9F24117C4B1FE649286651ECE"
24
+ "45B3DC2007CB8A163BF0598DA48361C55D39A69163FA8FD"
25
+ "24CF5F83655D23DCA3AD961C62F356208552BB9ED529077"
26
+ "096966D670C354E4ABC9804F1746C08CA18217C32905E46"
27
+ "2E36CE3BE39E772C180E86039B2783A2EC07A28FB5C55DF"
28
+ "06F4C52C9DE2BCBF6955817183995497CEA956AE515D226"
29
+ "1898FA051015728E5A8AACAA68FFFFFFFFFFFFFFFF"
30
+ )
31
+ P = int(_P_HEX, 16)
32
+ G = 2
33
+
34
+
35
+ @dataclass
36
+ class KexResult:
37
+ shared_secret: int # K
38
+ exchange_hash: bytes # H
39
+ host_key_blob: bytes # K_S, as received from the server
40
+ signature_blob: bytes # raw "string" signature blob from KEXDH_REPLY
41
+
42
+
43
+ def perform_diffie_hellman(
44
+ sock: socket.socket,
45
+ reader: SocketReader,
46
+ client_version: bytes,
47
+ server_version: bytes,
48
+ client_kexinit_payload: bytes,
49
+ server_kexinit_payload: bytes,
50
+ ) -> KexResult:
51
+ x = secrets.randbelow(P - 3) + 2 # private exponent, 2 <= x <= P-2
52
+ e = pow(G, x, P)
53
+
54
+ packet = encode_binary_packet(bytes([messages.SSH_MSG_KEXDH_INIT]) + encode_mpint(e))
55
+ sock.sendall(packet)
56
+
57
+ reply_payload = decode_binary_packet(reader.read_exact)
58
+ reply = Reader(reply_payload)
59
+ msg_type = reply.read_byte()
60
+ if msg_type != messages.SSH_MSG_KEXDH_REPLY:
61
+ raise ProtocolError(f"Expected SSH_MSG_KEXDH_REPLY ({messages.SSH_MSG_KEXDH_REPLY}), got {msg_type}")
62
+
63
+ host_key_blob = reply.read_string()
64
+ f = reply.read_mpint()
65
+ signature_blob = reply.read_string()
66
+
67
+ shared_secret = pow(f, x, P)
68
+ exchange_hash = compute_exchange_hash(
69
+ client_version=client_version,
70
+ server_version=server_version,
71
+ client_kexinit_payload=client_kexinit_payload,
72
+ server_kexinit_payload=server_kexinit_payload,
73
+ host_key_blob=host_key_blob,
74
+ e=e,
75
+ f=f,
76
+ shared_secret=shared_secret,
77
+ )
78
+ return KexResult(shared_secret, exchange_hash, host_key_blob, signature_blob)
79
+
80
+
81
+ def compute_exchange_hash(
82
+ *,
83
+ client_version: bytes,
84
+ server_version: bytes,
85
+ client_kexinit_payload: bytes,
86
+ server_kexinit_payload: bytes,
87
+ host_key_blob: bytes,
88
+ e: int,
89
+ f: int,
90
+ shared_secret: int,
91
+ ) -> bytes:
92
+ """H = SHA256(V_C || V_S || I_C || I_S || K_S || e || f || K), RFC 4253 §8."""
93
+ hash_input = b"".join(
94
+ [
95
+ encode_string(client_version),
96
+ encode_string(server_version),
97
+ encode_string(client_kexinit_payload),
98
+ encode_string(server_kexinit_payload),
99
+ encode_string(host_key_blob),
100
+ encode_mpint(e),
101
+ encode_mpint(f),
102
+ encode_mpint(shared_secret),
103
+ ]
104
+ )
105
+ return hashlib.sha256(hash_input).digest()