aws-ssh-utils 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.
@@ -0,0 +1 @@
1
+ __version__ = "0.1.0"
@@ -0,0 +1,173 @@
1
+ from dataclasses import dataclass
2
+ from typing import TYPE_CHECKING, MutableMapping
3
+
4
+ import click_spinner
5
+ import questionary
6
+
7
+ if TYPE_CHECKING:
8
+ from mypy_boto3_emr import EMRClient
9
+ from mypy_boto3_emr.type_defs import ClusterSummaryTypeDef, InstanceTypeDef
10
+
11
+
12
+ @dataclass
13
+ class IP:
14
+ private: str
15
+ public: str = None
16
+
17
+
18
+ def prompt_for_emr_cluster(
19
+ emr: "EMRClient",
20
+ prompt="Which cluster do you want to connect to?",
21
+ # Either a list of application names, or a dict[name, version]
22
+ applications: dict[str, str] | list[str] = None,
23
+ ) -> tuple[str, str]:
24
+ """Discovers the available EMR clusters and asks the user which to use.
25
+
26
+ Returns the cluster id of the cluster that the user selected.
27
+ """
28
+ # Discover available clusters
29
+ with click_spinner.spinner():
30
+ clusters = get_emr_clusters(emr, applications=applications)
31
+
32
+ # Prompt the user
33
+ cluster_name = questionary.select(
34
+ prompt,
35
+ choices=sorted(clusters),
36
+ ).unsafe_ask()
37
+ cluster_id, cluster_name = clusters[cluster_name]
38
+
39
+ return cluster_id, cluster_name
40
+
41
+
42
+ def prompt_for_emr_instance_group(
43
+ emr: "EMRClient",
44
+ cluster_id: str, prompt="Which instance group do you want to connect to?"
45
+ ) -> tuple[list[IP], str]:
46
+ """Discovers the available instance groups for the selected EMR cluster.
47
+ Note that instance group and instance fleet are both treated the same here.
48
+
49
+ Returns the list of instances in the selected group, as well as the group name
50
+ """
51
+ # Discover the available instances grouped by groupname
52
+ with click_spinner.spinner():
53
+ grouped_instances = get_emr_instance_ips(emr, cluster_id)
54
+
55
+ # Prompt the user
56
+ group_name = questionary.select(
57
+ prompt,
58
+ choices=sorted(grouped_instances),
59
+ ).unsafe_ask()
60
+
61
+ instances = grouped_instances[group_name]
62
+ return instances, group_name
63
+
64
+
65
+ def does_cluster_have_applications(
66
+ emr: "EMRClient",
67
+ cluster_id: str,
68
+ applications: dict[str, str] | list[str] = None,
69
+ ) -> bool:
70
+ cluster_details = emr.describe_cluster(ClusterId=cluster_id).get("Cluster")
71
+
72
+ if not cluster_details:
73
+ return False
74
+
75
+ cluster_applications = {a["Name"]: a["Version"] for a in cluster_details["Applications"]}
76
+
77
+ if isinstance(applications, MutableMapping):
78
+ # If a map, check that all the required applications are in the cluster and at the version
79
+ return all(v == cluster_applications.get(k) for k, v in applications.items())
80
+ else:
81
+ return all(k in cluster_applications for k in applications)
82
+
83
+
84
+ def get_emr_clusters(
85
+ emr: "EMRClient",
86
+ states: list[str] = None,
87
+ # Either a list of application names, or a dict[name, version]
88
+ applications: dict[str, str] | list[str] = None,
89
+ ) -> dict[str, tuple[str, str]]:
90
+ """Discover the available EMR cluster
91
+
92
+ Returns a dict where the key is the cluster "Name - ID - CreatedDT" and the value is the ID.
93
+ """
94
+ if states is None:
95
+ states = ["RUNNING", "WAITING"]
96
+
97
+ clusters = emr.list_clusters(
98
+ ClusterStates=states,
99
+ ).get("Clusters", [])
100
+
101
+ if applications:
102
+ clusters = [c for c in clusters if does_cluster_have_applications(emr, c["Id"], applications=applications)]
103
+
104
+ def get_display_name(c: "ClusterSummaryTypeDef") -> str:
105
+ created = c["Status"]["Timeline"]["CreationDateTime"].strftime("%Y-%m-%d %H:%M")
106
+ return f"{c['Name']} - {c['Id']} - {created}"
107
+
108
+ return {f"{get_display_name(c)}": (c["Id"], c["Name"]) for c in clusters}
109
+
110
+
111
+ def get_emr_groups(
112
+ emr: "EMRClient",
113
+ cluster_id: str,
114
+ ):
115
+ """Discover the EMR groups.
116
+
117
+ Returns the groups and the id_name to use for list_instances
118
+ """
119
+ cluster_details = emr.describe_cluster(ClusterId=cluster_id)
120
+ instance_collection_type = cluster_details["Cluster"].get("InstanceCollectionType")
121
+
122
+ if instance_collection_type == "INSTANCE_FLEET":
123
+ id_name = "InstanceFleetId"
124
+ groups = emr.list_instance_fleets(ClusterId=cluster_id)
125
+ groups = groups["InstanceFleets"]
126
+
127
+ elif instance_collection_type == "INSTANCE_GROUP":
128
+ id_name = "InstanceGroupId"
129
+ groups = emr.list_instance_groups(ClusterId=cluster_id)
130
+ groups = groups["InstanceGroups"]
131
+
132
+ return groups, id_name
133
+
134
+
135
+ def get_emr_instances(
136
+ emr: "EMRClient",
137
+ cluster_id: str,
138
+ ) -> dict[str, list["InstanceTypeDef"]]:
139
+ """Discover the running EMR instances per instance group"""
140
+ groups, id_name = get_emr_groups(emr, cluster_id)
141
+
142
+ groups = {g["Id"]: g["Name"] for g in groups}
143
+ grouped_instances = {g: [] for g in groups.values()}
144
+
145
+ instances = emr.list_instances(ClusterId=cluster_id, InstanceStates=["RUNNING"])
146
+ for instance in instances["Instances"]:
147
+ grouped_instances[groups[instance[id_name]]].append(instance)
148
+
149
+ return grouped_instances
150
+
151
+
152
+ def get_emr_instance_ips(
153
+ emr: "EMRClient",
154
+ cluster_id: str,
155
+ ) -> dict[str, list[IP]]:
156
+ '''Returns a dict of IPs per group'''
157
+ return {
158
+ k: [
159
+ IP(x["PrivateIpAddress"], x.get("PublicIpAddress"))
160
+ for x in v
161
+ ]
162
+ for k, v in get_emr_instances(emr, cluster_id).items()
163
+ }
164
+
165
+
166
+ def get_instance_key_name(
167
+ emr: "EMRClient",
168
+ cluster_id: str,
169
+ ) -> str:
170
+ cluster = emr.describe_cluster(ClusterId=cluster_id)
171
+
172
+ key_name = cluster["Cluster"]["Ec2InstanceAttributes"]["Ec2KeyName"]
173
+ return key_name
@@ -0,0 +1,280 @@
1
+ # Source: https://github.com/paramiko/paramiko/blob/main/demos/interactive.py
2
+
3
+ # Copyright (C) 2003-2007 Robey Pointer <robeypointer@gmail.com>
4
+ #
5
+ # This file is part of paramiko.
6
+ #
7
+ # Paramiko is free software; you can redistribute it and/or modify it under the
8
+ # terms of the GNU Lesser General Public License as published by the Free
9
+ # Software Foundation; either version 2.1 of the License, or (at your option)
10
+ # any later version.
11
+ #
12
+ # Paramiko is distributed in the hope that it will be useful, but WITHOUT ANY
13
+ # WARRANTY; without even the implied warranty of MERCHANTABILITY or FITNESS FOR
14
+ # A PARTICULAR PURPOSE. See the GNU Lesser General Public License for more
15
+ # details.
16
+ #
17
+ # You should have received a copy of the GNU Lesser General Public License
18
+ # along with Paramiko; if not, write to the Free Software Foundation, Inc.,
19
+ # 51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA.
20
+ import os
21
+ import socket
22
+ import sys
23
+
24
+ import paramiko
25
+
26
+ # https://invisible-island.net/xterm/ctlseqs/ctlseqs.html#h2-Bracketed-Paste-Mode
27
+ START_PASTE = "\x1B\x5B\x32\x30\x30\x7E" # ESC[200~
28
+ END_PASTE = "\x1B\x5B\x32\x30\x31\x7E" # ESC[201~
29
+
30
+
31
+ ALL_CODECS = [
32
+ 'ascii', 'big5', 'big5hkscs', 'cp037', 'cp273', 'cp424', 'cp437',
33
+ 'cp500', 'cp720', 'cp737', 'cp775', 'cp850', 'cp852', 'cp855', 'cp856', 'cp857',
34
+ 'cp858', 'cp860', 'cp861', 'cp862', 'cp863', 'cp864', 'cp865', 'cp866', 'cp869',
35
+ 'cp874', 'cp875', 'cp932', 'cp949', 'cp950', 'cp1006', 'cp1026', 'cp1125',
36
+ 'cp1140', 'cp1250', 'cp1251', 'cp1252', 'cp1253', 'cp1254', 'cp1255', 'cp1256',
37
+ 'cp1257', 'cp1258', 'euc_jp', 'euc_jis_2004', 'euc_jisx0213', 'euc_kr',
38
+ 'gb2312', 'gbk', 'gb18030', 'hz', 'iso2022_jp', 'iso2022_jp_1', 'iso2022_jp_2',
39
+ 'iso2022_jp_2004', 'iso2022_jp_3', 'iso2022_jp_ext', 'iso2022_kr', 'latin_1',
40
+ 'iso8859_2', 'iso8859_3', 'iso8859_4', 'iso8859_5', 'iso8859_6', 'iso8859_7',
41
+ 'iso8859_8', 'iso8859_9', 'iso8859_10', 'iso8859_11', 'iso8859_13',
42
+ 'iso8859_14', 'iso8859_15', 'iso8859_16', 'johab', 'koi8_r', 'koi8_t', 'koi8_u',
43
+ 'kz1048', 'mac_cyrillic', 'mac_greek', 'mac_iceland', 'mac_latin2', 'mac_roman',
44
+ 'mac_turkish', 'ptcp154', 'shift_jis', 'shift_jis_2004', 'shift_jisx0213',
45
+ 'utf_32', 'utf_32_be', 'utf_32_le', 'utf_16', 'utf_16_be', 'utf_16_le', 'utf_7',
46
+ 'utf_8', 'utf_8_sig'
47
+ ]
48
+
49
+ # windows does not have termios...
50
+ try:
51
+ import termios
52
+ import tty
53
+
54
+ has_termios = True
55
+ except ImportError:
56
+ has_termios = False
57
+
58
+ if os.getenv('TMUX'):
59
+ _TITLE_START = '\x1bk'
60
+ _TITLE_END = '\x1b\\'
61
+ else:
62
+ _TITLE_START = '\x1b]0;'
63
+ _TITLE_END = '\x07'
64
+
65
+
66
+ def is_int(val: str) -> bool:
67
+ try:
68
+ int(val)
69
+ return True
70
+ except Exception:
71
+ return False
72
+
73
+
74
+ def remove_title_change(data: str) -> str:
75
+ '''
76
+ Title change is a string starting with '\x1b]0;', and ending with '\x07'.
77
+ Remove all such substrings
78
+ '''
79
+ while True:
80
+ a = data.find(_TITLE_START)
81
+ b = data.find(_TITLE_END)
82
+
83
+ if a == -1 or b == -1:
84
+ break
85
+
86
+ data = data[:a] + data[b + len(_TITLE_END):]
87
+
88
+ return data
89
+
90
+
91
+ def interactive_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
92
+ if has_termios:
93
+ posix_shell(chan, allow_title_changes=allow_title_changes)
94
+ else:
95
+ windows_shell(chan, allow_title_changes=allow_title_changes)
96
+
97
+
98
+ def decode(chars: bytes) -> str:
99
+ '''
100
+ Decodes the bytes and handles encoding errors
101
+ `htop` scrolling for example uses cp037
102
+ '''
103
+ try:
104
+ # Attempt to decode the entire input as UTF-8
105
+ return chars.decode('utf-8')
106
+ except Exception:
107
+ # Attempt to decode each individual character
108
+ def decode_char(byte: bytes) -> str:
109
+ for codec in ALL_CODECS:
110
+ try:
111
+ c = byte.decode(codec)
112
+ return c
113
+ except Exception:
114
+ pass
115
+ else:
116
+ # Failed to decode character, return replacement character.
117
+ # https://www.fileformat.info/info/unicode/char/fffd/index.htm
118
+ return '\uFFFD'
119
+
120
+ return ''.join(decode_char(b for b in chars))
121
+
122
+
123
+ def posix_readkey() -> str:
124
+ """Get a keypress. If an escaped key is pressed, the full sequence is
125
+ read and returned.
126
+
127
+ Copied from readchar:
128
+ https://github.com/magmax/python-readchar/blob/master/readchar/_posix_read.py#L30
129
+ """
130
+
131
+ def read():
132
+ '''
133
+ Reads one character and handles encoding errors
134
+ `htop` scrolling for example uses cp037
135
+ '''
136
+ return decode(sys.stdin.buffer.raw.read(1))
137
+
138
+ c1 = read()
139
+
140
+ if c1 != "\x1B": # ESC
141
+ return c1
142
+
143
+ c2 = read()
144
+ if c2 not in "\x4F\x5B": # O[
145
+ return c1 + c2
146
+
147
+ c3 = read()
148
+ if c3 not in "\x31\x32\x33\x35\x36": # 12356
149
+ return c1 + c2 + c3
150
+
151
+ c4 = read()
152
+ if c4 not in "\x30\x31\x33\x34\x35\x37\x38\x39": # 01345789
153
+ return c1 + c2 + c3 + c4
154
+
155
+ c5 = read()
156
+ key = c1 + c2 + c3 + c4 + c5
157
+
158
+ # Bracketed Paste Mode: # https://invisible-island.net/xterm/ctlseqs/ctlseqs.html#h2-Bracketed-Paste-Mode
159
+ if key == START_PASTE[:-1] or key == END_PASTE[:-1]:
160
+ c6 = read()
161
+ return key + c6
162
+
163
+ return key
164
+
165
+
166
+ def windows_readkey() -> str:
167
+ """Reads the next keypress. If an escaped key is pressed, the full
168
+ sequence is read and returned.
169
+
170
+ Copied from readchar:
171
+ https://github.com/magmax/python-readchar/blob/master/readchar/_win_read.py#LL14C1-L30C24
172
+ """
173
+
174
+ ch = sys.stdin.read(1)
175
+
176
+ # if it is a normal character:
177
+ if ch not in "\x00\xe0":
178
+ return ch
179
+
180
+ # if it is a scpeal key, read second half:
181
+ ch2 = sys.stdin.read(1)
182
+
183
+ return "\x00" + ch2
184
+
185
+
186
+ def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True): # noqa: C901
187
+ import select
188
+
189
+ oldtty = termios.tcgetattr(sys.stdin)
190
+
191
+ # input_history = []
192
+ # output_history = []
193
+
194
+ try:
195
+ tty.setraw(sys.stdin.fileno())
196
+ tty.setcbreak(sys.stdin.fileno())
197
+ chan.settimeout(0.0)
198
+ while True:
199
+ r, w, e = select.select([chan, sys.stdin], [], [])
200
+ if chan in r:
201
+ try:
202
+ data = decode(chan.recv(1024))
203
+ if len(data) == 0:
204
+ sys.stdout.write("\r\n")
205
+ break
206
+
207
+ if not allow_title_changes:
208
+ data = remove_title_change(data)
209
+
210
+ # output_history.append(data)
211
+ sys.stdout.write(data)
212
+ sys.stdout.flush()
213
+ except socket.timeout:
214
+ pass
215
+ if sys.stdin in r:
216
+ key = posix_readkey()
217
+ # When pasting something, we need to read the entire pasted blob at once
218
+ # Otherwise it'll hang until the next key press.
219
+ # This has to do with how 'select.select' detects changes.
220
+ # A paste is a single event of many characters, so we must handle them all as one event
221
+ if key == START_PASTE:
222
+ # Start reading the pasted text
223
+ key = posix_readkey()
224
+ # Until we reach the end of the pasted text
225
+ while key != END_PASTE:
226
+ chan.send(key)
227
+ # input_history.append(key)
228
+ key = posix_readkey()
229
+ # We've exhausted the paste event, wait for next event
230
+ continue
231
+
232
+ if len(key) == 0:
233
+ break
234
+ chan.send(key)
235
+ # input_history.append(key)
236
+
237
+ finally:
238
+ termios.tcsetattr(sys.stdin, termios.TCSADRAIN, oldtty)
239
+
240
+ # Useful in debugging how control characters were send
241
+ # from pprint import pprint
242
+ # pprint(input_history)
243
+ # pprint(output_history)
244
+
245
+
246
+ # thanks to Mike Looijmans for this code
247
+ def windows_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
248
+ import threading
249
+
250
+ sys.stdout.write(
251
+ "Line-buffered terminal emulation. Press F6 or ^Z to send EOF.\r\n\r\n"
252
+ )
253
+
254
+ def writeall(sock):
255
+ while True:
256
+ data = sock.recv(256).decode()
257
+ if not data:
258
+ # Need user to input any character so we sys.stdin.read(1) completes and unblocks
259
+ sys.stdout.write("\r\n Connection closed. Press Enter to continue...\r\n")
260
+ sys.stdout.flush()
261
+ break
262
+
263
+ if not allow_title_changes:
264
+ data = remove_title_change(data)
265
+
266
+ sys.stdout.write(data)
267
+ sys.stdout.flush()
268
+
269
+ writer = threading.Thread(target=writeall, args=(chan,))
270
+ writer.start()
271
+
272
+ try:
273
+ while True:
274
+ d = windows_readkey()
275
+ if not d or chan.closed:
276
+ break
277
+ chan.send(d)
278
+ except EOFError:
279
+ # user hit ^Z or F6
280
+ pass
aws_ssh_utils/ssh.py ADDED
@@ -0,0 +1,493 @@
1
+ import logging
2
+ import os
3
+ import sys
4
+ import textwrap
5
+ from dataclasses import dataclass
6
+ from typing import TYPE_CHECKING
7
+
8
+ import boto3
9
+ import click
10
+ import click_spinner
11
+ import paramiko
12
+ import questionary
13
+ from botocore.exceptions import ClientError
14
+ from environs import Env
15
+ from loguru import logger
16
+
17
+ from .emr_utils import (
18
+ IP,
19
+ get_emr_instance_ips,
20
+ get_instance_key_name,
21
+ prompt_for_emr_cluster,
22
+ prompt_for_emr_instance_group,
23
+ )
24
+ from .interactive_ssh import interactive_shell
25
+
26
+ if TYPE_CHECKING:
27
+ from mypy_boto3_ec2.service_resource import Instance
28
+ from mypy_boto3_emr import EMRClient
29
+
30
+ Env().read_env() # Load .env file
31
+
32
+
33
+ class ShellError(Exception):
34
+ def __init__(self, message: str, exit_code: int = 1):
35
+ super().__init__(message, exit_code)
36
+ self.message = message
37
+ self.exit_code = exit_code
38
+
39
+
40
+ def set_terminal_title(title: str = ''):
41
+ if os.name == 'nt':
42
+ # Windows - CMD
43
+ os.system(f'title "{title}"')
44
+
45
+ # Windows - Powershell - But it seems that the scripts are always ran inside CMD anyway.
46
+ # os.system(f'$host.UI.RawUI.WindowTitle = {title}')
47
+ else:
48
+ # Unix
49
+ if os.getenv('TMUX'):
50
+ # Set the TMUX Window title
51
+ sys.stdout.write(f"\33k{title}\33")
52
+ else:
53
+ # Set the Gnome terminal title
54
+ sys.stdout.write(f"\33]0;{title}\a")
55
+ sys.stdout.flush()
56
+
57
+
58
+ @dataclass
59
+ class SelectedEMRInstance:
60
+ cluster_id: str
61
+ cluster_name: str
62
+ group_name: str
63
+ group_idx: str
64
+ ip: IP
65
+
66
+
67
+ @click.group()
68
+ @click.option('-ll/', '--long-log/--no-long-log', default=False, help='Enable long logging')
69
+ @click.option('--verbose/--no-verbose', default=False, help='Enable debug logging')
70
+ @click.option('--quiet/--no-quiet', default=False, help='Disable logging')
71
+ def cli(
72
+ long_log: bool = False,
73
+ verbose: bool = False,
74
+ quiet: bool = False,
75
+ **kwargs,
76
+ ):
77
+ format = (
78
+ "<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</green> | "
79
+ "<level>{level: <8}</level> | "
80
+ "<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>"
81
+ ) if long_log else "<level>{message}</level>"
82
+ level = "DEBUG" if verbose else 100 if quiet else "INFO"
83
+
84
+ logger.remove()
85
+ logger.add(sys.stdout, level=level, format=format, colorize=True)
86
+
87
+
88
+ @cli.command('ec2')
89
+ @click.option('-p', '--profile', default=None, help='Which AWS profile to use')
90
+ @click.option('-r', '--region', default=None, help='Which AWS region to use')
91
+ @click.option('-u', '--user', default=None, help='Which user to connect as')
92
+ @click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
93
+ @click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
94
+ def ec2_ssh(
95
+ profile: str = None,
96
+ region: str = None,
97
+ user: str = None,
98
+ private: bool = True,
99
+ key_file: str = None,
100
+ **kwargs,
101
+ ):
102
+ '''
103
+ Asks user which EC2 instance they want to connect to,
104
+ then opens an interactive SSH session to the instance
105
+
106
+ '''
107
+ try:
108
+ b3s = boto3.Session(profile_name=profile, region_name=region)
109
+
110
+ ip, user, key_file, instance_name = get_ec2_ssh_options(b3s=b3s, user=user, use_private_ip=private, key_file=key_file)
111
+ terminal_title = f'{user}@{instance_name}'
112
+
113
+ SSHShell(hostname=ip, username=user, key_filename=key_file, terminal_title=terminal_title).connect()
114
+ except ShellError as e:
115
+ logger.error(e.message)
116
+ exit(e.exit_code)
117
+ except ClientError as e:
118
+ if e.response.get('Error') and e.response['Error'].get('Code') == 'ExpiredTokenException':
119
+ logger.log(logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.')
120
+ exit(1)
121
+
122
+
123
+ @cli.command('emr')
124
+ @click.option('-p', '--profile', default=None, help='Which AWS profile to use')
125
+ @click.option('-r', '--region', default=None, help='Which AWS region to use')
126
+ @click.option('-u', '--user', default=None, help='Which user to connect as')
127
+ @click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
128
+ @click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
129
+ def emr_ssh(
130
+ profile: str = None,
131
+ region: str = None,
132
+ user: str = None,
133
+ private: bool = True,
134
+ key_file: str = None,
135
+ **kwargs
136
+ ):
137
+ '''
138
+ Asks user which Cluster and EC2 instance they want to connect to,
139
+ then opens an interactive SSH session to the instance
140
+ '''
141
+ try:
142
+ b3s = boto3.Session(profile_name=profile, region_name=region)
143
+ emr = b3s.client("emr")
144
+
145
+ emr_instance = get_emr_ssh_options(emr)
146
+
147
+ if user is None:
148
+ user = 'hadoop'
149
+
150
+ if key_file is None:
151
+ key_file = get_emr_ssh_key_file_for_cluster(emr, emr_instance.cluster_id)
152
+
153
+ if private:
154
+ ip = emr_instance.ip.private
155
+ else:
156
+ ip = emr_instance.ip.public
157
+ if ip is None:
158
+ raise ShellError('The selected instance does not have a public IP')
159
+
160
+ group_name = emr_instance.group_name.split(' ')[0]
161
+ postfix = f'[{emr_instance.group_idx}]' if emr_instance.group_idx is not None else ''
162
+ terminal_title = f'{emr_instance.cluster_name} - {group_name}{postfix}'
163
+
164
+ with SSHShell(hostname=ip, username=user, key_filename=key_file, terminal_title=terminal_title) as shell:
165
+ if group_name.lower().startswith('core'):
166
+ shell.send('sudo su\r\n')
167
+ shell.send('cd /var/log/hadoop-yarn/containers\r\n')
168
+
169
+ set_terminal_title()
170
+
171
+ except ShellError as e:
172
+ logger.error(e.message)
173
+ exit(e.exit_code)
174
+ except ClientError as e:
175
+ if e.response.get('Error') and e.response['Error'].get('Code') == 'ExpiredTokenException':
176
+ logger.log(logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.')
177
+ exit(1)
178
+
179
+
180
+ @cli.command('emr-all')
181
+ @click.option('-p', '--profile', default=None, help='Which AWS profile to use')
182
+ @click.option('-r', '--region', default=None, help='Which AWS region to use')
183
+ @click.option('-u', '--user', default=None, help='Which user to connect as')
184
+ @click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
185
+ @click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
186
+ def emr_ssh_all(
187
+ profile: str = None,
188
+ region: str = None,
189
+ user: str = None,
190
+ private: bool = True,
191
+ key_file: str = None,
192
+ **kwargs,
193
+ ):
194
+ '''
195
+ Asks user which Cluster and EC2 instance they want to connect to,
196
+ Then prints a tmux cli statement that will open a new session
197
+ with a window per ec2 instance with ssh shell already opened.
198
+ '''
199
+ try:
200
+ b3s = boto3.Session(profile_name=profile, region_name=region)
201
+ emr = b3s.client('emr')
202
+ cluster_id, cluster_name = prompt_for_emr_cluster(emr)
203
+ grouped_instances = {
204
+ k.split(" ")[0]: v
205
+ for k, v in get_emr_instance_ips(emr, cluster_id).items()
206
+ }
207
+
208
+ if user is None:
209
+ user = 'hadoop'
210
+
211
+ if key_file is None:
212
+ key_file = get_emr_ssh_key_file_for_cluster(emr, cluster_id)
213
+
214
+ window_order = ['Master', 'Primary', 'Core', 'Task']
215
+
216
+ def get_ip(ip: IP) -> str:
217
+ if private:
218
+ return ip.private
219
+ else:
220
+ if ip.public is None:
221
+ raise ShellError(f'The instance with private IP {ip.private} does not have a public IP')
222
+ return ip.public
223
+
224
+ window_cmds = [
225
+ # Add "|| $SHELL -i" so that if the ssh session fails, we still have a window so we can look at the error.
226
+ # without this, tmux will close the window.
227
+ (f'{group_name} - {instance_num}', f'ssh -i {key_file} {user}@{get_ip(instance_ip)} || $SHELL -i')
228
+ for group_name, instance_ips in sorted(grouped_instances.items(), key=lambda t: window_order.index(t[0]))
229
+ for instance_num, instance_ip in enumerate(instance_ips, start=1)
230
+ ]
231
+
232
+ session_name = cluster_name
233
+ first_window_name, first_cmd = window_cmds[0]
234
+ # Create a new tmux session with the EMR cluster name as the session name and open the first ssh connection
235
+ tmux(f'new-session -d -s "{session_name}" -n "{first_window_name}" "{first_cmd}"')
236
+ # Create a new tmux window and open the ssh connections
237
+ for window_cmd in window_cmds[1:]:
238
+ tmux(f'new-window -t "{session_name}:" -n "{window_cmd[0]}" "{window_cmd[1]}"')
239
+ # Now switch to the session's first window
240
+ tmux(f'switch-client -t "{session_name}:{first_window_name}"')
241
+
242
+ except ShellError as e:
243
+ logger.error(e.message)
244
+ exit(e.exit_code)
245
+ except ClientError as e:
246
+ if e.response.get('Error') and e.response['Error'].get('Code') == 'ExpiredTokenException':
247
+ logger.log(logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.')
248
+ exit(1)
249
+
250
+
251
+ @cli.command('test')
252
+ def test(**kwargs):
253
+ tmux('new-session -d -s "test" -n "window 1" "echo hello && $SHELL -i"')
254
+ tmux('new-window -n "window 2" -t "test:" "echo hello && $SHELL -i"')
255
+ tmux('switch-client -t "test:window 1"')
256
+
257
+
258
+ @dataclass
259
+ class SSHShell:
260
+ hostname: str
261
+ username: str
262
+ key_filename: str
263
+ terminal_title: str = None
264
+
265
+ def connect(self):
266
+ self._open()
267
+ self._launch()
268
+ self._close()
269
+
270
+ def _open(self):
271
+ logger.info(f'Opening SSH: ssh -i {self.key_filename} {self.username}@{self.hostname}')
272
+
273
+ terminal_size = os.get_terminal_size()
274
+
275
+ ssh_client = paramiko.SSHClient()
276
+ # Set hosts key path so we can save to it
277
+ host_key_path = os.path.expanduser('~/.ssh/known_hosts')
278
+ ssh_client.load_host_keys(host_key_path)
279
+ ssh_client.set_missing_host_key_policy(ConfirmAddPolicy())
280
+ try:
281
+ ssh_client.connect(hostname=self.hostname, username=self.username, key_filename=self.key_filename)
282
+ except paramiko.BadHostKeyException as e:
283
+ error_message = f'''
284
+ @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
285
+ @ WARNING: REMOTE HOST IDENTIFICATION HAS CHANGED! @
286
+ @@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
287
+ IT IS POSSIBLE THAT SOMEONE IS DOING SOMETHING NASTY!
288
+ Someone could be eavesdropping on you right now (man-in-the-middle attack)!
289
+ It is also possible that a host key has just been changed.
290
+ The fingerprint for the key sent by the remote host is
291
+ {e.key.fingerprint}
292
+ Expected key is
293
+ {e.expected_key.fingerprint}
294
+ Please contact your system administrator.
295
+ Add correct host key in {host_key_path} to get rid of this message.
296
+ Offending key in {host_key_path}
297
+ remove with:
298
+ ssh-keygen -f "{host_key_path}" -R "{self.hostname}"
299
+ Host key for {self.hostname} has changed and you have requested strict checking.
300
+ Host key verification failed.
301
+ '''
302
+ raise ShellError(textwrap.dedent(error_message), 255) from None
303
+
304
+ channel = ssh_client.get_transport().open_session()
305
+ channel.get_pty(term=os.getenv('TERM', 'xterm-256color'), width=terminal_size.columns, height=terminal_size.lines)
306
+ channel.invoke_shell()
307
+
308
+ self._ssh_client = ssh_client
309
+ self._channel = channel
310
+ if self.terminal_title:
311
+ set_terminal_title(self.terminal_title)
312
+
313
+ return channel
314
+
315
+ def _launch(self):
316
+ interactive_shell(self._channel, allow_title_changes=not self.terminal_title)
317
+
318
+ def _close(self):
319
+ self._ssh_client.close()
320
+ # Reset on connection close
321
+ if self.terminal_title:
322
+ set_terminal_title()
323
+
324
+ def __enter__(self):
325
+ return self._open()
326
+
327
+ def __exit__(self, exception_type, exception_value, traceback):
328
+ self._launch()
329
+ self._close()
330
+
331
+
332
+ def tmux(command: str):
333
+ os.system(f'tmux {command}')
334
+
335
+
336
+ def try_to_find_ssh_key_file(key_name: str) -> str:
337
+ '''Recursively iterate over the `~/.ssh/` folder to find the matching key'''
338
+ if key_name:
339
+ for dirpath, _, filenames in os.walk(os.path.expanduser('~/.ssh/')):
340
+ if filenames:
341
+ for f in filenames:
342
+ if f == key_name or os.path.splitext(f)[0] == key_name:
343
+ return os.path.join(dirpath, f)
344
+
345
+ return None
346
+
347
+
348
+ def get_ec2_image_name(instance) -> str:
349
+ try:
350
+ return instance.image.name
351
+ except AttributeError:
352
+ return None
353
+
354
+
355
+ def get_ec2_ssh_options(
356
+ b3s: boto3.Session,
357
+ user: str = None,
358
+ use_private_ip: bool = True,
359
+ key_file: str = None
360
+ ) -> tuple[str, str, str, str]:
361
+ instance = prompt_for_ec2_instance(b3s)
362
+ if key_file is None:
363
+ key_file = try_to_find_ssh_key_file(instance.key_name)
364
+
365
+ if user is None:
366
+ logger.info('No user specified, attempting to detect required user...')
367
+ image_name = get_ec2_image_name(instance)
368
+ if 'ubuntu' in image_name:
369
+ user = 'ubuntu'
370
+ else:
371
+ user = 'ec2-user'
372
+
373
+ if use_private_ip:
374
+ ip = instance.private_ip_address
375
+ elif instance.public_ip_address:
376
+ ip = instance.public_ip_address
377
+ else:
378
+ raise ShellError('Public IP was requested, but none was found!')
379
+
380
+ return ip, user, key_file, get_ec2_name(instance)
381
+
382
+
383
+ def get_emr_ssh_options(
384
+ emr: "EMRClient",
385
+ ) -> SelectedEMRInstance:
386
+ cluster_id, cluster_name = prompt_for_emr_cluster(emr)
387
+ instance_ips, group_name = prompt_for_emr_instance_group(emr, cluster_id)
388
+
389
+ instance_ips = sorted(instance_ips, key=lambda ips: ips.private)
390
+
391
+ instance_options = [
392
+ f'{ip.private} ({ip.public})' if ip.public else ip.private
393
+ for ip in instance_ips
394
+ ]
395
+ instance_ip = questionary.select(
396
+ 'Which instance do you want to connect to?',
397
+ choices=instance_options,
398
+ ).unsafe_ask()
399
+ group_idx = instance_options.index(instance_ip)
400
+
401
+ return SelectedEMRInstance(cluster_id, cluster_name, group_name, group_idx, instance_ips[group_idx])
402
+
403
+
404
+ def get_emr_ssh_key_file_for_cluster(
405
+ emr: "EMRClient",
406
+ cluster_id: str,
407
+ ) -> str:
408
+ with click_spinner.spinner():
409
+ key_name = get_instance_key_name(emr, cluster_id)
410
+ key_file = try_to_find_ssh_key_file(key_name)
411
+
412
+ if key_file is None:
413
+ should_continue = questionary.confirm(f'Could not find the ssh key {key_name}, would you like to continue?').unsafe_ask()
414
+ if not should_continue:
415
+ exit(1)
416
+
417
+ return key_file
418
+
419
+
420
+ def get_running_ec2_instances(b3s: boto3.Session):
421
+ '''Discover running instances'''
422
+ with click_spinner.spinner():
423
+ ec2 = b3s.resource('ec2')
424
+ running_instances = ec2.instances.filter(
425
+ Filters=[{'Name': 'instance-state-name', 'Values': ['running']}],
426
+ )
427
+
428
+ return running_instances
429
+
430
+
431
+ def prompt_for_ec2_instance(
432
+ b3s: boto3.Session,
433
+ prompt: str = 'Which EC2 instance do you want to connect to?',
434
+ ) -> "Instance":
435
+ '''
436
+ Discovers the runnin EC2 clusters and asks the user which to use.
437
+
438
+ Returns the EC2 object that the user selected.
439
+ '''
440
+ running_instances = get_running_ec2_instances(b3s)
441
+
442
+ name_contains = questionary.text(
443
+ 'Provide a name filter or leave blank to show all',
444
+ ).unsafe_ask().lower()
445
+
446
+ with click_spinner.spinner():
447
+ grouped_by_name = {
448
+ get_ec2_name(instance): instance
449
+ for instance in running_instances
450
+ if not name_contains or name_contains in get_ec2_name(instance).lower()
451
+ }
452
+
453
+ if not grouped_by_name:
454
+ raise ShellError('No matching EC2 instances found!')
455
+
456
+ # Prompt the user
457
+ ec2_name = questionary.select(
458
+ prompt,
459
+ choices=sorted(grouped_by_name),
460
+ ).unsafe_ask()
461
+ instance = grouped_by_name[ec2_name]
462
+
463
+ return instance
464
+
465
+
466
+ def get_ec2_name(ec2_instance: "Instance") -> str:
467
+ '''Takes in the boto3 EC2 resource instance object'''
468
+ for tag in ec2_instance.tags:
469
+ if tag['Key'] == 'Name':
470
+ return tag['Value']
471
+
472
+ return ec2_instance.instance_id
473
+
474
+
475
+ class ConfirmAddPolicy(paramiko.client.MissingHostKeyPolicy):
476
+ """
477
+ Policy for automatically adding the hostname and new host key to the
478
+ local `.HostKeys` object, and saving it. This is used by `.SSHClient`.
479
+ """
480
+
481
+ def missing_host_key(self, client, hostname, key):
482
+ logger.warning(f"Unknown {key.get_name()} host key for {hostname}: {key.fingerprint}")
483
+ should_add = questionary.confirm("Continue and add host key?").unsafe_ask()
484
+ if should_add:
485
+ client._host_keys.add(hostname, key.get_name(), key)
486
+ if client._host_keys_filename is not None:
487
+ client.save_host_keys(client._host_keys_filename)
488
+ logger.info('Added host key')
489
+ else:
490
+ logger.warning('Failed to add host key. No host key file defined!')
491
+
492
+ else:
493
+ raise paramiko.SSHException(f"Server {hostname!r} not found in known_hosts")
@@ -0,0 +1,113 @@
1
+ Metadata-Version: 2.4
2
+ Name: aws-ssh-utils
3
+ Version: 0.1.0
4
+ Summary: Easy AWS SSHing
5
+ Project-URL: Homepage, https://github.com/mvanderlee/aws_ssh_utils
6
+ Author-email: Michiel Vanderlee <jmt.vanderlee@gmail.com>
7
+ License-Expression: MIT
8
+ Classifier: Intended Audience :: Developers
9
+ Classifier: Intended Audience :: Information Technology
10
+ Classifier: Intended Audience :: System Administrators
11
+ Classifier: License :: OSI Approved :: MIT License
12
+ Classifier: Operating System :: OS Independent
13
+ Classifier: Programming Language :: Python
14
+ Classifier: Programming Language :: Python :: 3 :: Only
15
+ Classifier: Programming Language :: Python :: 3.11
16
+ Classifier: Programming Language :: Python :: 3.12
17
+ Classifier: Programming Language :: Python :: 3.13
18
+ Classifier: Topic :: Software Development
19
+ Classifier: Topic :: Software Development :: Libraries
20
+ Classifier: Topic :: Software Development :: Libraries :: Python Modules
21
+ Classifier: Typing :: Typed
22
+ Requires-Python: >=3.11
23
+ Requires-Dist: boto3
24
+ Requires-Dist: click
25
+ Requires-Dist: click-spinner
26
+ Requires-Dist: coloredlogs
27
+ Requires-Dist: environs
28
+ Requires-Dist: loguru
29
+ Requires-Dist: paramiko
30
+ Requires-Dist: questionary
31
+ Provides-Extra: dev
32
+ Requires-Dist: boto3-stubs[ec2,emr]; extra == 'dev'
33
+ Provides-Extra: publish
34
+ Requires-Dist: hatch>=1.7.0; extra == 'publish'
35
+ Provides-Extra: test
36
+ Description-Content-Type: text/markdown
37
+
38
+ # AWS Auth
39
+
40
+ ```shell
41
+ pip install aws-auth-utils
42
+
43
+ aws configure --profile mfa-source
44
+
45
+ aws_auth mfa
46
+ ```
47
+
48
+ The commands use [click](https://click.palletsprojects.com/en/stable/) for argument parsing and if required arguments are missing it will prompt you.
49
+
50
+ To authenticate using your MFA token you will need to have a profile configured using regular an AWS Access Key.
51
+
52
+ We will use that and your MFA token to generate an authorized session profile.
53
+ By default we will try to use the `mfa-source` and create the `default` profile.
54
+
55
+ If you only have a single MFA device set up, it will use that automatically. If you have multiple, it will the first one.
56
+
57
+ ## MFA
58
+
59
+ ```shell
60
+ $ aws_auth mfa --help
61
+ Usage: aws_auth mfa [OPTIONS]
62
+
63
+ Options:
64
+ -a, --mfa-arn TEXT The identification number of the MFA device that
65
+ is associated with the IAM user. i.e.:
66
+ "arn:aws:iam::123456789012:mfa/tony.stark". You
67
+ can find this on the IAM page.
68
+ -c, --code TEXT The code generated by your MFA device.
69
+ -d, --duration INTEGER The duration, in seconds, of the session.
70
+ -sp, --source-profile TEXT What AWS profile to get the session token with.
71
+ -tp, --target-profile TEXT What AWS profile to store the credentials under.
72
+ -v, --verbose BOOLEAN
73
+ --help Show this message and exit.
74
+ ```
75
+
76
+ ## Assume Role
77
+
78
+ The assume role is useful for multi-org environments where you want to impersonate a role in a child organization.
79
+ If you access multiple organizations I recommend you set up aliases.
80
+
81
+ ```shell
82
+ aws_auth assume \
83
+ --role-arn arn:aws:iam::123456789012:role/OrganizationAccountAccessRole \
84
+ --session-name child_org \
85
+ --target-profile child_session
86
+ ```
87
+
88
+ ```shell
89
+ $ aws_auth assume --help
90
+ Usage: aws_auth assume [OPTIONS]
91
+
92
+ Get MFA authenticated and assumed role session credentials and save them to
93
+ the aws credentials file
94
+
95
+ If you have multiple accounts you'd like to switch between, I recommend
96
+ setting up aliases that call this script with predefined arguments.
97
+
98
+ Options:
99
+ -r, --role-arn TEXT The Arn of the Role to assume.
100
+ -n, --session-name TEXT The identifier for the assumed role session.
101
+ -a, --mfa-arn TEXT The identification number of the MFA device that
102
+ is associated with the IAM user. i.e.:
103
+ "arn:aws:iam::123456789012:mfa/tony.stark". You
104
+ can find this on the IAM page.
105
+ -c, --code TEXT The code generated by your MFA device.
106
+ -d, --duration INTEGER The duration, in seconds, of the session.
107
+ (defaults to 4 hours)
108
+ -sp, --source-profile TEXT What AWS profile to get the session token with.
109
+ -tp, --target-profile TEXT What AWS profile to store the credentials under.
110
+ -v, --verbose BOOLEAN
111
+ --help Show this message and exit.
112
+ ```
113
+
@@ -0,0 +1,8 @@
1
+ aws_ssh_utils/__init__.py,sha256=kUR5RAFc7HCeiqdlX36dZOHkUI5wI6V_43RpEcD8b-0,22
2
+ aws_ssh_utils/emr_utils.py,sha256=RyAVP2IGXIMfo1pqBmfupFUJruLyxiwEcKYNhA7LK7w,5410
3
+ aws_ssh_utils/interactive_ssh.py,sha256=gcBgDOnIExQmg5Fh6JpAfcA5KeezGjw-Qr2b6rE0LNM,8975
4
+ aws_ssh_utils/ssh.py,sha256=zsqPHMJU3WxhqqhQMTXCzrxmi6IfUliTzgOLyQjTc9A,17375
5
+ aws_ssh_utils-0.1.0.dist-info/METADATA,sha256=Ug2CW31EHg_iU2p0h3LcNsMff3Mq5XenzWo-X2-uLys,4321
6
+ aws_ssh_utils-0.1.0.dist-info/WHEEL,sha256=qtCwoSJWgHk21S1Kb4ihdzI2rlJ1ZKaIurTj_ngOhyQ,87
7
+ aws_ssh_utils-0.1.0.dist-info/entry_points.txt,sha256=H3d3Sn9YhxItIEzmMmnrOIqfvwGuLHEYbhiI64itUV4,50
8
+ aws_ssh_utils-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.27.0
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,2 @@
1
+ [console_scripts]
2
+ aws_ssh = aws_ssh_utils.ssh:cli