aws-ssh-utils 0.1.1__tar.gz → 0.2.1__tar.gz
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.
- aws_ssh_utils-0.2.1/LICENSE +21 -0
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/PKG-INFO +8 -2
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/README.md +4 -0
- aws_ssh_utils-0.2.1/aws_ssh_utils/__init__.py +1 -0
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/aws_ssh_utils/emr_utils.py +50 -34
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/aws_ssh_utils/interactive_ssh.py +33 -31
- aws_ssh_utils-0.2.1/aws_ssh_utils/py.typed +0 -0
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/aws_ssh_utils/ssh.py +214 -89
- aws_ssh_utils-0.2.1/pyproject.toml +161 -0
- aws_ssh_utils-0.1.1/aws_ssh_utils/__init__.py +0 -1
- aws_ssh_utils-0.1.1/pyproject.toml +0 -102
- {aws_ssh_utils-0.1.1 → aws_ssh_utils-0.2.1}/.gitignore +0 -0
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2025 mvanderlee
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -1,10 +1,11 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: aws-ssh-utils
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.2.1
|
|
4
4
|
Summary: Easy AWS SSHing
|
|
5
|
-
Project-URL: Homepage, https://github.com/mvanderlee/
|
|
5
|
+
Project-URL: Homepage, https://github.com/mvanderlee/aws-ssh-utils
|
|
6
6
|
Author-email: Michiel Vanderlee <jmt.vanderlee@gmail.com>
|
|
7
7
|
License-Expression: MIT
|
|
8
|
+
License-File: LICENSE
|
|
8
9
|
Classifier: Intended Audience :: Developers
|
|
9
10
|
Classifier: Intended Audience :: Information Technology
|
|
10
11
|
Classifier: Intended Audience :: System Administrators
|
|
@@ -28,6 +29,7 @@ Requires-Dist: environs
|
|
|
28
29
|
Requires-Dist: loguru
|
|
29
30
|
Requires-Dist: paramiko
|
|
30
31
|
Requires-Dist: questionary
|
|
32
|
+
Requires-Dist: typing-extensions
|
|
31
33
|
Provides-Extra: dev
|
|
32
34
|
Requires-Dist: boto3-stubs[ec2,emr]; extra == 'dev'
|
|
33
35
|
Provides-Extra: publish
|
|
@@ -37,6 +39,10 @@ Description-Content-Type: text/markdown
|
|
|
37
39
|
|
|
38
40
|
# AWS SSH Utils
|
|
39
41
|
|
|
42
|
+
[](https://pypi.org/project/aws-ssh-utils/)
|
|
43
|
+
[](#)
|
|
44
|
+
[](https://pypi.org/project/aws-ssh-utils/)
|
|
45
|
+
|
|
40
46
|
```shell
|
|
41
47
|
pip install aws-ssh-utils
|
|
42
48
|
|
|
@@ -1,5 +1,9 @@
|
|
|
1
1
|
# AWS SSH Utils
|
|
2
2
|
|
|
3
|
+
[](https://pypi.org/project/aws-ssh-utils/)
|
|
4
|
+
[](#)
|
|
5
|
+
[](https://pypi.org/project/aws-ssh-utils/)
|
|
6
|
+
|
|
3
7
|
```shell
|
|
4
8
|
pip install aws-ssh-utils
|
|
5
9
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.2.1"
|
|
@@ -1,25 +1,27 @@
|
|
|
1
|
+
from collections.abc import MutableMapping, Sequence
|
|
1
2
|
from dataclasses import dataclass
|
|
2
|
-
from typing import TYPE_CHECKING,
|
|
3
|
+
from typing import TYPE_CHECKING, Literal
|
|
3
4
|
|
|
4
5
|
import click_spinner
|
|
5
6
|
import questionary
|
|
6
7
|
|
|
7
8
|
if TYPE_CHECKING:
|
|
8
9
|
from mypy_boto3_emr import EMRClient
|
|
9
|
-
from mypy_boto3_emr.
|
|
10
|
+
from mypy_boto3_emr.literals import ClusterStateType
|
|
11
|
+
from mypy_boto3_emr.type_defs import ClusterSummaryTypeDef, InstanceFleetTypeDef, InstanceGroupTypeDef, InstanceTypeDef
|
|
10
12
|
|
|
11
13
|
|
|
12
14
|
@dataclass
|
|
13
15
|
class IP:
|
|
14
|
-
private: str
|
|
15
|
-
public: str = None
|
|
16
|
+
private: str | None = None
|
|
17
|
+
public: str | None = None
|
|
16
18
|
|
|
17
19
|
|
|
18
20
|
def prompt_for_emr_cluster(
|
|
19
21
|
emr: "EMRClient",
|
|
20
|
-
prompt="Which cluster do you want to connect to?",
|
|
22
|
+
prompt: str = "Which cluster do you want to connect to?",
|
|
21
23
|
# Either a list of application names, or a dict[name, version]
|
|
22
|
-
applications: dict[str, str] | list[str] = None,
|
|
24
|
+
applications: dict[str, str] | list[str] | None = None,
|
|
23
25
|
) -> tuple[str, str]:
|
|
24
26
|
"""Discovers the available EMR clusters and asks the user which to use.
|
|
25
27
|
|
|
@@ -41,7 +43,8 @@ def prompt_for_emr_cluster(
|
|
|
41
43
|
|
|
42
44
|
def prompt_for_emr_instance_group(
|
|
43
45
|
emr: "EMRClient",
|
|
44
|
-
cluster_id: str,
|
|
46
|
+
cluster_id: str,
|
|
47
|
+
prompt: str = "Which instance group do you want to connect to?",
|
|
45
48
|
) -> tuple[list[IP], str]:
|
|
46
49
|
"""Discovers the available instance groups for the selected EMR cluster.
|
|
47
50
|
Note that instance group and instance fleet are both treated the same here.
|
|
@@ -64,16 +67,20 @@ def prompt_for_emr_instance_group(
|
|
|
64
67
|
|
|
65
68
|
def does_cluster_have_applications(
|
|
66
69
|
emr: "EMRClient",
|
|
67
|
-
cluster_id: str,
|
|
68
|
-
applications: dict[str, str] | list[str] = None,
|
|
70
|
+
cluster_id: str | None = None,
|
|
71
|
+
applications: dict[str, str] | list[str] | None = None,
|
|
69
72
|
) -> bool:
|
|
70
|
-
|
|
73
|
+
if applications is None or cluster_id is None:
|
|
74
|
+
return True
|
|
71
75
|
|
|
76
|
+
cluster_details = emr.describe_cluster(ClusterId=cluster_id).get("Cluster")
|
|
72
77
|
if not cluster_details:
|
|
73
78
|
return False
|
|
74
79
|
|
|
75
|
-
cluster_applications = {
|
|
76
|
-
|
|
80
|
+
cluster_applications = {
|
|
81
|
+
a.get("Name"): a.get("Version")
|
|
82
|
+
for a in cluster_details.get("Applications", [])
|
|
83
|
+
}
|
|
77
84
|
if isinstance(applications, MutableMapping):
|
|
78
85
|
# If a map, check that all the required applications are in the cluster and at the version
|
|
79
86
|
return all(v == cluster_applications.get(k) for k, v in applications.items())
|
|
@@ -83,9 +90,9 @@ def does_cluster_have_applications(
|
|
|
83
90
|
|
|
84
91
|
def get_emr_clusters(
|
|
85
92
|
emr: "EMRClient",
|
|
86
|
-
states:
|
|
93
|
+
states: Sequence["ClusterStateType"] | None = None,
|
|
87
94
|
# Either a list of application names, or a dict[name, version]
|
|
88
|
-
applications: dict[str, str] | list[str] = None,
|
|
95
|
+
applications: dict[str, str] | list[str] | None = None,
|
|
89
96
|
) -> dict[str, tuple[str, str]]:
|
|
90
97
|
"""Discover the available EMR cluster
|
|
91
98
|
|
|
@@ -99,25 +106,37 @@ def get_emr_clusters(
|
|
|
99
106
|
).get("Clusters", [])
|
|
100
107
|
|
|
101
108
|
if applications:
|
|
102
|
-
clusters = [
|
|
109
|
+
clusters = [
|
|
110
|
+
c for c in clusters
|
|
111
|
+
if does_cluster_have_applications(emr, c.get("Id"), applications=applications)
|
|
112
|
+
]
|
|
103
113
|
|
|
104
114
|
def get_display_name(c: "ClusterSummaryTypeDef") -> str:
|
|
105
|
-
created = c
|
|
106
|
-
|
|
115
|
+
created = c.get("Status", {}).get(
|
|
116
|
+
"Timeline", {},
|
|
117
|
+
).get("CreationDateTime")
|
|
118
|
+
if created:
|
|
119
|
+
created = created.strftime("%Y-%m-%d %H:%M")
|
|
120
|
+
return f"{c.get('Name')} - {c.get('Id')} - {created}"
|
|
107
121
|
|
|
108
|
-
return {
|
|
122
|
+
return {
|
|
123
|
+
f"{get_display_name(c)}": (c.get("Id", ""), c.get("Name", ""))
|
|
124
|
+
for c in clusters
|
|
125
|
+
}
|
|
109
126
|
|
|
110
127
|
|
|
111
128
|
def get_emr_groups(
|
|
112
129
|
emr: "EMRClient",
|
|
113
130
|
cluster_id: str,
|
|
114
|
-
):
|
|
131
|
+
) -> tuple[list["InstanceFleetTypeDef"] | list["InstanceGroupTypeDef"], Literal['InstanceFleetId', 'InstanceGroupId']]:
|
|
115
132
|
"""Discover the EMR groups.
|
|
116
133
|
|
|
117
134
|
Returns the groups and the id_name to use for list_instances
|
|
118
135
|
"""
|
|
119
136
|
cluster_details = emr.describe_cluster(ClusterId=cluster_id)
|
|
120
|
-
instance_collection_type = cluster_details["Cluster"].get(
|
|
137
|
+
instance_collection_type = cluster_details["Cluster"].get(
|
|
138
|
+
"InstanceCollectionType",
|
|
139
|
+
)
|
|
121
140
|
|
|
122
141
|
if instance_collection_type == "INSTANCE_FLEET":
|
|
123
142
|
id_name = "InstanceFleetId"
|
|
@@ -129,6 +148,11 @@ def get_emr_groups(
|
|
|
129
148
|
groups = emr.list_instance_groups(ClusterId=cluster_id)
|
|
130
149
|
groups = groups["InstanceGroups"]
|
|
131
150
|
|
|
151
|
+
else:
|
|
152
|
+
raise ValueError(
|
|
153
|
+
f"Unknown instance collection type: {instance_collection_type}",
|
|
154
|
+
)
|
|
155
|
+
|
|
132
156
|
return groups, id_name
|
|
133
157
|
|
|
134
158
|
|
|
@@ -139,12 +163,14 @@ def get_emr_instances(
|
|
|
139
163
|
"""Discover the running EMR instances per instance group"""
|
|
140
164
|
groups, id_name = get_emr_groups(emr, cluster_id)
|
|
141
165
|
|
|
142
|
-
groups = {g
|
|
166
|
+
groups = {g.get("Id", ""): g.get("Name", "") for g in groups}
|
|
143
167
|
grouped_instances = {g: [] for g in groups.values()}
|
|
144
168
|
|
|
145
|
-
instances = emr.list_instances(
|
|
169
|
+
instances = emr.list_instances(
|
|
170
|
+
ClusterId=cluster_id, InstanceStates=["RUNNING"],
|
|
171
|
+
)
|
|
146
172
|
for instance in instances["Instances"]:
|
|
147
|
-
grouped_instances[groups[instance
|
|
173
|
+
grouped_instances[groups[instance.get(id_name, "")]].append(instance)
|
|
148
174
|
|
|
149
175
|
return grouped_instances
|
|
150
176
|
|
|
@@ -156,18 +182,8 @@ def get_emr_instance_ips(
|
|
|
156
182
|
'''Returns a dict of IPs per group'''
|
|
157
183
|
return {
|
|
158
184
|
k: [
|
|
159
|
-
IP(x
|
|
185
|
+
IP(x.get("PrivateIpAddress"), x.get("PublicIpAddress"))
|
|
160
186
|
for x in v
|
|
161
187
|
]
|
|
162
188
|
for k, v in get_emr_instances(emr, cluster_id).items()
|
|
163
189
|
}
|
|
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
|
|
@@ -20,8 +20,10 @@
|
|
|
20
20
|
import os
|
|
21
21
|
import socket
|
|
22
22
|
import sys
|
|
23
|
+
from importlib.util import find_spec
|
|
23
24
|
|
|
24
25
|
import paramiko
|
|
26
|
+
from loguru import logger
|
|
25
27
|
|
|
26
28
|
# https://invisible-island.net/xterm/ctlseqs/ctlseqs.html#h2-Bracketed-Paste-Mode
|
|
27
29
|
START_PASTE = "\x1B\x5B\x32\x30\x30\x7E" # ESC[200~
|
|
@@ -43,24 +45,15 @@ ALL_CODECS = [
|
|
|
43
45
|
'kz1048', 'mac_cyrillic', 'mac_greek', 'mac_iceland', 'mac_latin2', 'mac_roman',
|
|
44
46
|
'mac_turkish', 'ptcp154', 'shift_jis', 'shift_jis_2004', 'shift_jisx0213',
|
|
45
47
|
'utf_32', 'utf_32_be', 'utf_32_le', 'utf_16', 'utf_16_be', 'utf_16_le', 'utf_7',
|
|
46
|
-
'utf_8', 'utf_8_sig'
|
|
48
|
+
'utf_8', 'utf_8_sig',
|
|
47
49
|
]
|
|
48
50
|
|
|
49
51
|
# windows does not have termios...
|
|
50
|
-
|
|
51
|
-
import termios
|
|
52
|
-
import tty
|
|
52
|
+
has_termios = find_spec("termios") is not None and find_spec("tty") is not None
|
|
53
53
|
|
|
54
|
-
has_termios = True
|
|
55
|
-
except ImportError:
|
|
56
|
-
has_termios = False
|
|
57
54
|
|
|
58
|
-
if os.getenv('TMUX')
|
|
59
|
-
|
|
60
|
-
_TITLE_END = '\x1b\\'
|
|
61
|
-
else:
|
|
62
|
-
_TITLE_START = '\x1b]0;'
|
|
63
|
-
_TITLE_END = '\x07'
|
|
55
|
+
_TITLE_START = '\x1bk' if os.getenv('TMUX') else '\x1b]0;'
|
|
56
|
+
_TITLE_END = '\x1b\\' if os.getenv('TMUX') else '\x07'
|
|
64
57
|
|
|
65
58
|
|
|
66
59
|
def is_int(val: str) -> bool:
|
|
@@ -105,19 +98,21 @@ def decode(chars: bytes) -> str:
|
|
|
105
98
|
return chars.decode('utf-8')
|
|
106
99
|
except Exception:
|
|
107
100
|
# Attempt to decode each individual character
|
|
108
|
-
def decode_char(byte:
|
|
101
|
+
def decode_char(byte: int) -> str:
|
|
109
102
|
for codec in ALL_CODECS:
|
|
110
103
|
try:
|
|
111
|
-
c = byte.decode(codec)
|
|
104
|
+
c = byte.to_bytes(1, 'big').decode(codec)
|
|
112
105
|
return c
|
|
113
106
|
except Exception:
|
|
114
|
-
|
|
107
|
+
logger.debug(
|
|
108
|
+
f'Failed to decode character {byte} as {codec}',
|
|
109
|
+
)
|
|
115
110
|
else:
|
|
116
111
|
# Failed to decode character, return replacement character.
|
|
117
112
|
# https://www.fileformat.info/info/unicode/char/fffd/index.htm
|
|
118
113
|
return '\uFFFD'
|
|
119
114
|
|
|
120
|
-
return ''.join(decode_char(b for b in chars)
|
|
115
|
+
return ''.join(decode_char(b) for b in chars)
|
|
121
116
|
|
|
122
117
|
|
|
123
118
|
def posix_readkey() -> str:
|
|
@@ -133,7 +128,7 @@ def posix_readkey() -> str:
|
|
|
133
128
|
Reads one character and handles encoding errors
|
|
134
129
|
`htop` scrolling for example uses cp037
|
|
135
130
|
'''
|
|
136
|
-
return decode(sys.stdin.buffer.raw.read(1))
|
|
131
|
+
return decode(sys.stdin.buffer.raw.read(1)) # pyright: ignore[reportAttributeAccessIssue, reportUnknownArgumentType]
|
|
137
132
|
|
|
138
133
|
c1 = read()
|
|
139
134
|
|
|
@@ -183,20 +178,25 @@ def windows_readkey() -> str:
|
|
|
183
178
|
return "\x00" + ch2
|
|
184
179
|
|
|
185
180
|
|
|
186
|
-
def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
|
|
181
|
+
def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
|
|
182
|
+
if not has_termios:
|
|
183
|
+
raise RuntimeError("Termios is not available on this system")
|
|
184
|
+
|
|
187
185
|
import select
|
|
186
|
+
import termios
|
|
187
|
+
import tty
|
|
188
188
|
|
|
189
|
-
oldtty = termios.tcgetattr(sys.stdin)
|
|
189
|
+
oldtty = termios.tcgetattr(sys.stdin) # pyright: ignore[reportAttributeAccessIssue]
|
|
190
190
|
|
|
191
191
|
# input_history = []
|
|
192
192
|
# output_history = []
|
|
193
193
|
|
|
194
194
|
try:
|
|
195
|
-
tty.setraw(sys.stdin.fileno())
|
|
196
|
-
tty.setcbreak(sys.stdin.fileno())
|
|
195
|
+
tty.setraw(sys.stdin.fileno()) # pyright: ignore[reportAttributeAccessIssue]
|
|
196
|
+
tty.setcbreak(sys.stdin.fileno()) # pyright: ignore[reportAttributeAccessIssue]
|
|
197
197
|
chan.settimeout(0.0)
|
|
198
198
|
while True:
|
|
199
|
-
r,
|
|
199
|
+
r, _, _ = select.select([chan, sys.stdin], [], [])
|
|
200
200
|
if chan in r:
|
|
201
201
|
try:
|
|
202
202
|
data = decode(chan.recv(1024))
|
|
@@ -210,7 +210,7 @@ def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True): # no
|
|
|
210
210
|
# output_history.append(data)
|
|
211
211
|
sys.stdout.write(data)
|
|
212
212
|
sys.stdout.flush()
|
|
213
|
-
except
|
|
213
|
+
except TimeoutError:
|
|
214
214
|
pass
|
|
215
215
|
if sys.stdin in r:
|
|
216
216
|
key = posix_readkey()
|
|
@@ -223,7 +223,7 @@ def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True): # no
|
|
|
223
223
|
key = posix_readkey()
|
|
224
224
|
# Until we reach the end of the pasted text
|
|
225
225
|
while key != END_PASTE:
|
|
226
|
-
chan.send(key)
|
|
226
|
+
chan.send(key.encode())
|
|
227
227
|
# input_history.append(key)
|
|
228
228
|
key = posix_readkey()
|
|
229
229
|
# We've exhausted the paste event, wait for next event
|
|
@@ -231,11 +231,11 @@ def posix_shell(chan: paramiko.Channel, allow_title_changes: bool = True): # no
|
|
|
231
231
|
|
|
232
232
|
if len(key) == 0:
|
|
233
233
|
break
|
|
234
|
-
chan.send(key)
|
|
234
|
+
chan.send(key.encode())
|
|
235
235
|
# input_history.append(key)
|
|
236
236
|
|
|
237
237
|
finally:
|
|
238
|
-
termios.tcsetattr(sys.stdin, termios.TCSADRAIN, oldtty)
|
|
238
|
+
termios.tcsetattr(sys.stdin, termios.TCSADRAIN, oldtty) # pyright: ignore[reportAttributeAccessIssue]
|
|
239
239
|
|
|
240
240
|
# Useful in debugging how control characters were send
|
|
241
241
|
# from pprint import pprint
|
|
@@ -248,15 +248,17 @@ def windows_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
|
|
|
248
248
|
import threading
|
|
249
249
|
|
|
250
250
|
sys.stdout.write(
|
|
251
|
-
"Line-buffered terminal emulation. Press F6 or ^Z to send EOF.\r\n\r\n"
|
|
251
|
+
"Line-buffered terminal emulation. Press F6 or ^Z to send EOF.\r\n\r\n",
|
|
252
252
|
)
|
|
253
253
|
|
|
254
|
-
def writeall(sock):
|
|
254
|
+
def writeall(sock: socket.socket):
|
|
255
255
|
while True:
|
|
256
256
|
data = sock.recv(256).decode()
|
|
257
257
|
if not data:
|
|
258
258
|
# Need user to input any character so we sys.stdin.read(1) completes and unblocks
|
|
259
|
-
sys.stdout.write(
|
|
259
|
+
sys.stdout.write(
|
|
260
|
+
"\r\n Connection closed. Press Enter to continue...\r\n",
|
|
261
|
+
)
|
|
260
262
|
sys.stdout.flush()
|
|
261
263
|
break
|
|
262
264
|
|
|
@@ -274,7 +276,7 @@ def windows_shell(chan: paramiko.Channel, allow_title_changes: bool = True):
|
|
|
274
276
|
d = windows_readkey()
|
|
275
277
|
if not d or chan.closed:
|
|
276
278
|
break
|
|
277
|
-
chan.send(d)
|
|
279
|
+
chan.send(d.encode())
|
|
278
280
|
except EOFError:
|
|
279
281
|
# user hit ^Z or F6
|
|
280
282
|
pass
|
|
File without changes
|
|
@@ -1,23 +1,28 @@
|
|
|
1
|
+
import datetime as dt
|
|
1
2
|
import logging
|
|
2
3
|
import os
|
|
4
|
+
import subprocess
|
|
3
5
|
import sys
|
|
4
6
|
import textwrap
|
|
5
|
-
from dataclasses import dataclass
|
|
6
|
-
from
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from hashlib import sha256
|
|
9
|
+
from types import TracebackType
|
|
10
|
+
from typing import TYPE_CHECKING, Any
|
|
7
11
|
|
|
8
12
|
import boto3
|
|
9
13
|
import click
|
|
10
14
|
import click_spinner
|
|
11
15
|
import paramiko
|
|
16
|
+
import paramiko.client
|
|
12
17
|
import questionary
|
|
13
18
|
from botocore.exceptions import ClientError
|
|
14
19
|
from environs import Env
|
|
15
20
|
from loguru import logger
|
|
21
|
+
from typing_extensions import override
|
|
16
22
|
|
|
17
23
|
from .emr_utils import (
|
|
18
24
|
IP,
|
|
19
25
|
get_emr_instance_ips,
|
|
20
|
-
get_instance_key_name,
|
|
21
26
|
prompt_for_emr_cluster,
|
|
22
27
|
prompt_for_emr_instance_group,
|
|
23
28
|
)
|
|
@@ -28,19 +33,24 @@ if TYPE_CHECKING:
|
|
|
28
33
|
from mypy_boto3_emr import EMRClient
|
|
29
34
|
|
|
30
35
|
Env().read_env() # Load .env file
|
|
36
|
+
# CSV of issuer,client
|
|
37
|
+
OPKSSH_PROVIDER_TAG = 'opkssh_provider'
|
|
38
|
+
# EMR doesn't support comma's in tags.
|
|
39
|
+
OPKSSH_ISSUER_TAG = 'opkssh_issuer'
|
|
40
|
+
OPKSSH_CLIENT_TAG = 'opkssh_client'
|
|
31
41
|
|
|
32
42
|
|
|
33
43
|
class ShellError(Exception):
|
|
34
44
|
def __init__(self, message: str, exit_code: int = 1):
|
|
35
45
|
super().__init__(message, exit_code)
|
|
36
|
-
self.message = message
|
|
37
|
-
self.exit_code = exit_code
|
|
46
|
+
self.message: str = message
|
|
47
|
+
self.exit_code: int = exit_code
|
|
38
48
|
|
|
39
49
|
|
|
40
50
|
def set_terminal_title(title: str = ''):
|
|
41
51
|
if os.name == 'nt':
|
|
42
52
|
# Windows - CMD
|
|
43
|
-
os.system(f'title "{title}"')
|
|
53
|
+
os.system(f'title "{title}"') # noqa: S605
|
|
44
54
|
|
|
45
55
|
# Windows - Powershell - But it seems that the scripts are always ran inside CMD anyway.
|
|
46
56
|
# os.system(f'$host.UI.RawUI.WindowTitle = {title}')
|
|
@@ -60,7 +70,7 @@ class SelectedEMRInstance:
|
|
|
60
70
|
cluster_id: str
|
|
61
71
|
cluster_name: str
|
|
62
72
|
group_name: str
|
|
63
|
-
group_idx:
|
|
73
|
+
group_idx: int
|
|
64
74
|
ip: IP
|
|
65
75
|
|
|
66
76
|
|
|
@@ -72,9 +82,9 @@ def cli(
|
|
|
72
82
|
long_log: bool = False,
|
|
73
83
|
verbose: bool = False,
|
|
74
84
|
quiet: bool = False,
|
|
75
|
-
**kwargs,
|
|
85
|
+
**kwargs: Any,
|
|
76
86
|
):
|
|
77
|
-
|
|
87
|
+
log_format = (
|
|
78
88
|
"<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</green> | "
|
|
79
89
|
"<level>{level: <8}</level> | "
|
|
80
90
|
"<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>"
|
|
@@ -82,7 +92,7 @@ def cli(
|
|
|
82
92
|
level = "DEBUG" if verbose else 100 if quiet else "INFO"
|
|
83
93
|
|
|
84
94
|
logger.remove()
|
|
85
|
-
logger.add(sys.stdout, level=level, format=
|
|
95
|
+
logger.add(sys.stdout, level=level, format=log_format, colorize=True)
|
|
86
96
|
|
|
87
97
|
|
|
88
98
|
@cli.command('ec2')
|
|
@@ -92,12 +102,12 @@ def cli(
|
|
|
92
102
|
@click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
|
|
93
103
|
@click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
|
|
94
104
|
def ec2_ssh(
|
|
95
|
-
profile: str = None,
|
|
96
|
-
region: str = None,
|
|
97
|
-
user: str = None,
|
|
105
|
+
profile: str | None = None,
|
|
106
|
+
region: str | None = None,
|
|
107
|
+
user: str | None = None,
|
|
98
108
|
private: bool = True,
|
|
99
|
-
key_file: str = None,
|
|
100
|
-
**kwargs,
|
|
109
|
+
key_file: str | None = None,
|
|
110
|
+
**kwargs: Any,
|
|
101
111
|
):
|
|
102
112
|
'''
|
|
103
113
|
Asks user which EC2 instance they want to connect to,
|
|
@@ -107,16 +117,23 @@ def ec2_ssh(
|
|
|
107
117
|
try:
|
|
108
118
|
b3s = boto3.Session(profile_name=profile, region_name=region)
|
|
109
119
|
|
|
110
|
-
ip, user, key_file, instance_name = get_ec2_ssh_options(
|
|
120
|
+
ip, user, key_file, instance_name = get_ec2_ssh_options(
|
|
121
|
+
b3s=b3s, user=user, use_private_ip=private, key_file=key_file,
|
|
122
|
+
)
|
|
111
123
|
terminal_title = f'{user}@{instance_name}'
|
|
112
124
|
|
|
113
|
-
SSHShell(
|
|
125
|
+
SSHShell(
|
|
126
|
+
hostname=ip, username=user, key_filename=key_file,
|
|
127
|
+
terminal_title=terminal_title,
|
|
128
|
+
).connect()
|
|
114
129
|
except ShellError as e:
|
|
115
130
|
logger.error(e.message)
|
|
116
131
|
exit(e.exit_code)
|
|
117
132
|
except ClientError as e:
|
|
118
|
-
if e.response.get('Error') and e.response
|
|
119
|
-
logger.log(
|
|
133
|
+
if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
|
|
134
|
+
logger.log(
|
|
135
|
+
logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
|
|
136
|
+
)
|
|
120
137
|
exit(1)
|
|
121
138
|
|
|
122
139
|
|
|
@@ -127,12 +144,12 @@ def ec2_ssh(
|
|
|
127
144
|
@click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
|
|
128
145
|
@click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
|
|
129
146
|
def emr_ssh(
|
|
130
|
-
profile: str = None,
|
|
131
|
-
region: str = None,
|
|
132
|
-
user: str = None,
|
|
147
|
+
profile: str | None = None,
|
|
148
|
+
region: str | None = None,
|
|
149
|
+
user: str | None = None,
|
|
133
150
|
private: bool = True,
|
|
134
|
-
key_file: str = None,
|
|
135
|
-
**kwargs
|
|
151
|
+
key_file: str | None = None,
|
|
152
|
+
**kwargs: Any,
|
|
136
153
|
):
|
|
137
154
|
'''
|
|
138
155
|
Asks user which Cluster and EC2 instance they want to connect to,
|
|
@@ -148,23 +165,31 @@ def emr_ssh(
|
|
|
148
165
|
user = 'hadoop'
|
|
149
166
|
|
|
150
167
|
if key_file is None:
|
|
151
|
-
key_file = get_emr_ssh_key_file_for_cluster(
|
|
168
|
+
key_file = get_emr_ssh_key_file_for_cluster(
|
|
169
|
+
emr, emr_instance.cluster_id,
|
|
170
|
+
)
|
|
152
171
|
|
|
153
172
|
if private:
|
|
154
173
|
ip = emr_instance.ip.private
|
|
174
|
+
if ip is None:
|
|
175
|
+
raise ShellError(
|
|
176
|
+
'The selected instance does not have a private IP',
|
|
177
|
+
)
|
|
155
178
|
else:
|
|
156
179
|
ip = emr_instance.ip.public
|
|
157
180
|
if ip is None:
|
|
158
|
-
raise ShellError(
|
|
181
|
+
raise ShellError(
|
|
182
|
+
'The selected instance does not have a public IP',
|
|
183
|
+
)
|
|
159
184
|
|
|
160
185
|
group_name = emr_instance.group_name.split(' ')[0]
|
|
161
|
-
postfix = f'[{emr_instance.group_idx}]'
|
|
186
|
+
postfix = f'[{emr_instance.group_idx}]' or ''
|
|
162
187
|
terminal_title = f'{emr_instance.cluster_name} - {group_name}{postfix}'
|
|
163
188
|
|
|
164
189
|
with SSHShell(hostname=ip, username=user, key_filename=key_file, terminal_title=terminal_title) as shell:
|
|
165
190
|
if group_name.lower().startswith('core'):
|
|
166
|
-
shell.send('sudo su\r\n')
|
|
167
|
-
shell.send('cd /var/log/hadoop-yarn/containers\r\n')
|
|
191
|
+
shell.send(b'sudo su\r\n')
|
|
192
|
+
shell.send(b'cd /var/log/hadoop-yarn/containers\r\n')
|
|
168
193
|
|
|
169
194
|
set_terminal_title()
|
|
170
195
|
|
|
@@ -172,8 +197,10 @@ def emr_ssh(
|
|
|
172
197
|
logger.error(e.message)
|
|
173
198
|
exit(e.exit_code)
|
|
174
199
|
except ClientError as e:
|
|
175
|
-
if e.response.get('Error') and e.response
|
|
176
|
-
logger.log(
|
|
200
|
+
if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
|
|
201
|
+
logger.log(
|
|
202
|
+
logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
|
|
203
|
+
)
|
|
177
204
|
exit(1)
|
|
178
205
|
|
|
179
206
|
|
|
@@ -184,12 +211,12 @@ def emr_ssh(
|
|
|
184
211
|
@click.option('--private/--public', default=True, help="Connect to the instance's private or public IP")
|
|
185
212
|
@click.option('-k/', '--key-file', default=None, help="Which key file to use to connect")
|
|
186
213
|
def emr_ssh_all(
|
|
187
|
-
profile: str = None,
|
|
188
|
-
region: str = None,
|
|
189
|
-
user: str = None,
|
|
214
|
+
profile: str | None = None,
|
|
215
|
+
region: str | None = None,
|
|
216
|
+
user: str | None = None,
|
|
190
217
|
private: bool = True,
|
|
191
|
-
key_file: str = None,
|
|
192
|
-
**kwargs,
|
|
218
|
+
key_file: str | None = None,
|
|
219
|
+
**kwargs: Any,
|
|
193
220
|
):
|
|
194
221
|
'''
|
|
195
222
|
Asks user which Cluster and EC2 instance they want to connect to,
|
|
@@ -213,18 +240,23 @@ def emr_ssh_all(
|
|
|
213
240
|
|
|
214
241
|
window_order = ['Master', 'Primary', 'Core', 'Task']
|
|
215
242
|
|
|
216
|
-
def get_ip(ip: IP) -> str:
|
|
243
|
+
def get_ip(ip: IP) -> str | None:
|
|
217
244
|
if private:
|
|
218
245
|
return ip.private
|
|
219
246
|
else:
|
|
220
247
|
if ip.public is None:
|
|
221
|
-
raise ShellError(
|
|
248
|
+
raise ShellError(
|
|
249
|
+
f'The instance with private IP {ip.private} does not have a public IP',
|
|
250
|
+
)
|
|
222
251
|
return ip.public
|
|
223
252
|
|
|
224
253
|
window_cmds = [
|
|
225
254
|
# Add "|| $SHELL -i" so that if the ssh session fails, we still have a window so we can look at the error.
|
|
226
255
|
# without this, tmux will close the window.
|
|
227
|
-
(
|
|
256
|
+
(
|
|
257
|
+
f'{group_name} - {instance_num}',
|
|
258
|
+
f'ssh -i {key_file} {user}@{get_ip(instance_ip)} || $SHELL -i',
|
|
259
|
+
)
|
|
228
260
|
for group_name, instance_ips in sorted(grouped_instances.items(), key=lambda t: window_order.index(t[0]))
|
|
229
261
|
for instance_num, instance_ip in enumerate(instance_ips, start=1)
|
|
230
262
|
]
|
|
@@ -232,10 +264,14 @@ def emr_ssh_all(
|
|
|
232
264
|
session_name = cluster_name
|
|
233
265
|
first_window_name, first_cmd = window_cmds[0]
|
|
234
266
|
# Create a new tmux session with the EMR cluster name as the session name and open the first ssh connection
|
|
235
|
-
tmux(
|
|
267
|
+
tmux(
|
|
268
|
+
f'new-session -d -s "{session_name}" -n "{first_window_name}" "{first_cmd}"',
|
|
269
|
+
)
|
|
236
270
|
# Create a new tmux window and open the ssh connections
|
|
237
271
|
for window_cmd in window_cmds[1:]:
|
|
238
|
-
tmux(
|
|
272
|
+
tmux(
|
|
273
|
+
f'new-window -t "{session_name}:" -n "{window_cmd[0]}" "{window_cmd[1]}"',
|
|
274
|
+
)
|
|
239
275
|
# Now switch to the session's first window
|
|
240
276
|
tmux(f'switch-client -t "{session_name}:{first_window_name}"')
|
|
241
277
|
|
|
@@ -243,32 +279,43 @@ def emr_ssh_all(
|
|
|
243
279
|
logger.error(e.message)
|
|
244
280
|
exit(e.exit_code)
|
|
245
281
|
except ClientError as e:
|
|
246
|
-
if e.response.get('Error') and e.response
|
|
247
|
-
logger.log(
|
|
282
|
+
if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
|
|
283
|
+
logger.log(
|
|
284
|
+
logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
|
|
285
|
+
)
|
|
248
286
|
exit(1)
|
|
249
287
|
|
|
250
288
|
|
|
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
289
|
@dataclass
|
|
259
290
|
class SSHShell:
|
|
260
291
|
hostname: str
|
|
261
292
|
username: str
|
|
262
|
-
key_filename: str
|
|
263
|
-
terminal_title: str = None
|
|
293
|
+
key_filename: str | None = None
|
|
294
|
+
terminal_title: str | None = None
|
|
295
|
+
|
|
296
|
+
private_key: paramiko.PKey = field(init=False)
|
|
297
|
+
_ssh_client: paramiko.SSHClient = field(init=False)
|
|
298
|
+
_channel: paramiko.Channel = field(init=False)
|
|
299
|
+
|
|
300
|
+
def __post_init__(self):
|
|
301
|
+
if self.key_filename is not None:
|
|
302
|
+
if not os.path.isfile(self.key_filename):
|
|
303
|
+
raise ValueError(f'File {self.key_filename} does not exist')
|
|
304
|
+
|
|
305
|
+
self.private_key = paramiko.PKey.from_path(self.key_filename)
|
|
306
|
+
public_key_path = f'{self.key_filename}.pub'
|
|
307
|
+
if os.path.isfile(public_key_path):
|
|
308
|
+
self.private_key.load_certificate(public_key_path)
|
|
264
309
|
|
|
265
310
|
def connect(self):
|
|
266
311
|
self._open()
|
|
267
312
|
self._launch()
|
|
268
313
|
self._close()
|
|
269
314
|
|
|
270
|
-
def _open(self):
|
|
271
|
-
logger.info(
|
|
315
|
+
def _open(self) -> paramiko.Channel:
|
|
316
|
+
logger.info(
|
|
317
|
+
f'Opening SSH: ssh -i {self.key_filename} {self.username}@{self.hostname}',
|
|
318
|
+
)
|
|
272
319
|
|
|
273
320
|
terminal_size = os.get_terminal_size()
|
|
274
321
|
|
|
@@ -278,7 +325,11 @@ class SSHShell:
|
|
|
278
325
|
ssh_client.load_host_keys(host_key_path)
|
|
279
326
|
ssh_client.set_missing_host_key_policy(ConfirmAddPolicy())
|
|
280
327
|
try:
|
|
281
|
-
ssh_client.connect(
|
|
328
|
+
ssh_client.connect(
|
|
329
|
+
hostname=self.hostname,
|
|
330
|
+
username=self.username,
|
|
331
|
+
pkey=self.private_key,
|
|
332
|
+
)
|
|
282
333
|
except paramiko.BadHostKeyException as e:
|
|
283
334
|
error_message = f'''
|
|
284
335
|
@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@@
|
|
@@ -301,8 +352,11 @@ class SSHShell:
|
|
|
301
352
|
'''
|
|
302
353
|
raise ShellError(textwrap.dedent(error_message), 255) from None
|
|
303
354
|
|
|
304
|
-
channel = ssh_client.get_transport().open_session()
|
|
305
|
-
channel.get_pty(
|
|
355
|
+
channel = ssh_client.get_transport().open_session() # pyright: ignore[reportOptionalMemberAccess]
|
|
356
|
+
channel.get_pty(
|
|
357
|
+
term=os.getenv('TERM', 'xterm-256color'),
|
|
358
|
+
width=terminal_size.columns, height=terminal_size.lines,
|
|
359
|
+
)
|
|
306
360
|
channel.invoke_shell()
|
|
307
361
|
|
|
308
362
|
self._ssh_client = ssh_client
|
|
@@ -313,7 +367,9 @@ class SSHShell:
|
|
|
313
367
|
return channel
|
|
314
368
|
|
|
315
369
|
def _launch(self):
|
|
316
|
-
interactive_shell(
|
|
370
|
+
interactive_shell(
|
|
371
|
+
self._channel, allow_title_changes=not self.terminal_title,
|
|
372
|
+
)
|
|
317
373
|
|
|
318
374
|
def _close(self):
|
|
319
375
|
self._ssh_client.close()
|
|
@@ -324,16 +380,16 @@ class SSHShell:
|
|
|
324
380
|
def __enter__(self):
|
|
325
381
|
return self._open()
|
|
326
382
|
|
|
327
|
-
def __exit__(self, exception_type, exception_value, traceback):
|
|
383
|
+
def __exit__(self, exception_type: type[BaseException] | None, exception_value: BaseException | None, traceback: TracebackType | None):
|
|
328
384
|
self._launch()
|
|
329
385
|
self._close()
|
|
330
386
|
|
|
331
387
|
|
|
332
388
|
def tmux(command: str):
|
|
333
|
-
os.system(f'tmux {command}')
|
|
389
|
+
os.system(f'tmux {command}') # noqa: S605
|
|
334
390
|
|
|
335
391
|
|
|
336
|
-
def try_to_find_ssh_key_file(key_name: str) -> str:
|
|
392
|
+
def try_to_find_ssh_key_file(key_name: str) -> str | None:
|
|
337
393
|
'''Recursively iterate over the `~/.ssh/` folder to find the matching key'''
|
|
338
394
|
if key_name:
|
|
339
395
|
for dirpath, _, filenames in os.walk(os.path.expanduser('~/.ssh/')):
|
|
@@ -345,7 +401,7 @@ def try_to_find_ssh_key_file(key_name: str) -> str:
|
|
|
345
401
|
return None
|
|
346
402
|
|
|
347
403
|
|
|
348
|
-
def get_ec2_image_name(instance) -> str:
|
|
404
|
+
def get_ec2_image_name(instance: "Instance") -> str | None:
|
|
349
405
|
try:
|
|
350
406
|
return instance.image.name
|
|
351
407
|
except AttributeError:
|
|
@@ -354,21 +410,34 @@ def get_ec2_image_name(instance) -> str:
|
|
|
354
410
|
|
|
355
411
|
def get_ec2_ssh_options(
|
|
356
412
|
b3s: boto3.Session,
|
|
357
|
-
user: str = None,
|
|
413
|
+
user: str | None = None,
|
|
358
414
|
use_private_ip: bool = True,
|
|
359
|
-
key_file: str = None
|
|
360
|
-
) -> tuple[str, str, str, str]:
|
|
415
|
+
key_file: str | None = None,
|
|
416
|
+
) -> tuple[str, str, str | None, str]:
|
|
361
417
|
instance = prompt_for_ec2_instance(b3s)
|
|
418
|
+
instance_tags = {
|
|
419
|
+
tag.get('Key', '').lower(): tag.get('Value', '')
|
|
420
|
+
for tag in instance.tags
|
|
421
|
+
}
|
|
422
|
+
|
|
362
423
|
if key_file is None:
|
|
363
|
-
|
|
424
|
+
# Support https://github.com/openpubkey/opkssh via tags
|
|
425
|
+
if OPKSSH_PROVIDER_TAG in instance_tags:
|
|
426
|
+
key_file = get_opkssh_key_file(instance_tags[OPKSSH_PROVIDER_TAG])
|
|
427
|
+
elif OPKSSH_ISSUER_TAG in instance_tags and OPKSSH_CLIENT_TAG in instance_tags:
|
|
428
|
+
key_file = get_opkssh_key_file(
|
|
429
|
+
','.join([
|
|
430
|
+
instance_tags[OPKSSH_ISSUER_TAG],
|
|
431
|
+
instance_tags[OPKSSH_CLIENT_TAG],
|
|
432
|
+
]),
|
|
433
|
+
)
|
|
434
|
+
else:
|
|
435
|
+
key_file = try_to_find_ssh_key_file(instance.key_name)
|
|
364
436
|
|
|
365
437
|
if user is None:
|
|
366
438
|
logger.info('No user specified, attempting to detect required user...')
|
|
367
439
|
image_name = get_ec2_image_name(instance)
|
|
368
|
-
if 'ubuntu' in image_name
|
|
369
|
-
user = 'ubuntu'
|
|
370
|
-
else:
|
|
371
|
-
user = 'ec2-user'
|
|
440
|
+
user = 'ubuntu' if image_name and 'ubuntu' in image_name.lower() else 'ec2-user'
|
|
372
441
|
|
|
373
442
|
if use_private_ip:
|
|
374
443
|
ip = instance.private_ip_address
|
|
@@ -386,10 +455,10 @@ def get_emr_ssh_options(
|
|
|
386
455
|
cluster_id, cluster_name = prompt_for_emr_cluster(emr)
|
|
387
456
|
instance_ips, group_name = prompt_for_emr_instance_group(emr, cluster_id)
|
|
388
457
|
|
|
389
|
-
instance_ips = sorted(instance_ips, key=lambda ips: ips.private)
|
|
458
|
+
instance_ips = sorted(instance_ips, key=lambda ips: ips.private or "")
|
|
390
459
|
|
|
391
460
|
instance_options = [
|
|
392
|
-
f'{ip.private} ({ip.public})' if ip.public else ip.private
|
|
461
|
+
f'{ip.private} ({ip.public})' if ip.public else f'{ip.private}'
|
|
393
462
|
for ip in instance_ips
|
|
394
463
|
]
|
|
395
464
|
instance_ip = questionary.select(
|
|
@@ -404,15 +473,38 @@ def get_emr_ssh_options(
|
|
|
404
473
|
def get_emr_ssh_key_file_for_cluster(
|
|
405
474
|
emr: "EMRClient",
|
|
406
475
|
cluster_id: str,
|
|
407
|
-
) -> str:
|
|
476
|
+
) -> str | None:
|
|
408
477
|
with click_spinner.spinner():
|
|
409
|
-
|
|
410
|
-
|
|
478
|
+
cluster = emr.describe_cluster(ClusterId=cluster_id)
|
|
479
|
+
cluster_tags = {
|
|
480
|
+
tag.get('Key', '').lower(): tag.get('Value', '')
|
|
481
|
+
for tag in cluster['Cluster'].get('Tags', [])
|
|
482
|
+
}
|
|
411
483
|
|
|
412
|
-
|
|
413
|
-
|
|
414
|
-
if
|
|
415
|
-
|
|
484
|
+
key_file = key_name = None
|
|
485
|
+
# Support https://github.com/openpubkey/opkssh via tags
|
|
486
|
+
if OPKSSH_PROVIDER_TAG in cluster_tags:
|
|
487
|
+
key_file = get_opkssh_key_file(cluster_tags[OPKSSH_PROVIDER_TAG])
|
|
488
|
+
elif OPKSSH_ISSUER_TAG in cluster_tags and OPKSSH_CLIENT_TAG in cluster_tags:
|
|
489
|
+
key_file = get_opkssh_key_file(
|
|
490
|
+
','.join([
|
|
491
|
+
cluster_tags[OPKSSH_ISSUER_TAG],
|
|
492
|
+
cluster_tags[OPKSSH_CLIENT_TAG],
|
|
493
|
+
]),
|
|
494
|
+
)
|
|
495
|
+
|
|
496
|
+
if key_file is None:
|
|
497
|
+
key_name = cluster["Cluster"].get(
|
|
498
|
+
"Ec2InstanceAttributes", {},
|
|
499
|
+
).get("Ec2KeyName", "")
|
|
500
|
+
key_file = try_to_find_ssh_key_file(key_name)
|
|
501
|
+
|
|
502
|
+
if key_file is None:
|
|
503
|
+
should_continue = questionary.confirm(
|
|
504
|
+
f'Could not find the ssh key {key_name}, would you like to continue?',
|
|
505
|
+
).unsafe_ask()
|
|
506
|
+
if not should_continue:
|
|
507
|
+
exit(1)
|
|
416
508
|
|
|
417
509
|
return key_file
|
|
418
510
|
|
|
@@ -466,28 +558,61 @@ def prompt_for_ec2_instance(
|
|
|
466
558
|
def get_ec2_name(ec2_instance: "Instance") -> str:
|
|
467
559
|
'''Takes in the boto3 EC2 resource instance object'''
|
|
468
560
|
for tag in ec2_instance.tags:
|
|
469
|
-
if tag
|
|
470
|
-
return tag
|
|
561
|
+
if tag.get('Key') == 'Name':
|
|
562
|
+
return tag.get('Value', '')
|
|
471
563
|
|
|
472
564
|
return ec2_instance.instance_id
|
|
473
565
|
|
|
474
566
|
|
|
567
|
+
def get_opkssh_key_file(provider: str) -> str:
|
|
568
|
+
opkssh_key_file = os.path.join(
|
|
569
|
+
os.path.expanduser(
|
|
570
|
+
'~/.ssh/',
|
|
571
|
+
), f'opkssh_{sha256(provider.encode()).hexdigest()}',
|
|
572
|
+
)
|
|
573
|
+
|
|
574
|
+
if os.path.exists(opkssh_key_file) and os.path.getmtime(opkssh_key_file) > dt.datetime.now().timestamp() - 86400:
|
|
575
|
+
logger.info('Found existing opkssh key')
|
|
576
|
+
else:
|
|
577
|
+
logger.info('Detected opkssh, logging in')
|
|
578
|
+
cmd = f"opkssh login --provider '{provider}' -i '{opkssh_key_file}'"
|
|
579
|
+
process = subprocess.Popen(cmd, shell=True, stdout=subprocess.PIPE) # noqa: S602
|
|
580
|
+
process.wait()
|
|
581
|
+
if process.returncode != 0:
|
|
582
|
+
logger.error('Failed to authenticate using opkssh')
|
|
583
|
+
sys.exit(process.returncode)
|
|
584
|
+
|
|
585
|
+
return opkssh_key_file
|
|
586
|
+
|
|
587
|
+
|
|
475
588
|
class ConfirmAddPolicy(paramiko.client.MissingHostKeyPolicy):
|
|
476
589
|
"""
|
|
477
590
|
Policy for automatically adding the hostname and new host key to the
|
|
478
591
|
local `.HostKeys` object, and saving it. This is used by `.SSHClient`.
|
|
479
592
|
"""
|
|
480
593
|
|
|
481
|
-
|
|
482
|
-
|
|
483
|
-
|
|
594
|
+
@override
|
|
595
|
+
def missing_host_key(self, client: paramiko.SSHClient, hostname: str, key: paramiko.PKey):
|
|
596
|
+
logger.warning(
|
|
597
|
+
f"Unknown {key.get_name()} host key for {hostname}: {key.fingerprint}",
|
|
598
|
+
)
|
|
599
|
+
should_add = questionary.confirm(
|
|
600
|
+
"Continue and add host key?",
|
|
601
|
+
).unsafe_ask()
|
|
484
602
|
if should_add:
|
|
485
|
-
client.
|
|
486
|
-
|
|
487
|
-
client.
|
|
603
|
+
client.get_host_keys().add(hostname, key.get_name(), key)
|
|
604
|
+
host_key_filename: str | None = (
|
|
605
|
+
client._host_keys_filename # pyright: ignore[reportAttributeAccessIssue]
|
|
606
|
+
)
|
|
607
|
+
if host_key_filename is not None:
|
|
608
|
+
client.save_host_keys(host_key_filename)
|
|
488
609
|
logger.info('Added host key')
|
|
489
610
|
else:
|
|
490
|
-
logger.warning(
|
|
611
|
+
logger.warning(
|
|
612
|
+
'Failed to add host key. No host key file defined!',
|
|
613
|
+
)
|
|
491
614
|
|
|
492
615
|
else:
|
|
493
|
-
raise paramiko.SSHException(
|
|
616
|
+
raise paramiko.SSHException(
|
|
617
|
+
f"Server {hostname!r} not found in known_hosts",
|
|
618
|
+
)
|
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["hatchling"]
|
|
3
|
+
build-backend = "hatchling.build"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "aws-ssh-utils"
|
|
7
|
+
description = "Easy AWS SSHing"
|
|
8
|
+
readme = "README.md"
|
|
9
|
+
requires-python = ">=3.11"
|
|
10
|
+
license = "MIT"
|
|
11
|
+
authors = [
|
|
12
|
+
{ name = "Michiel Vanderlee", email = "jmt.vanderlee@gmail.com" },
|
|
13
|
+
]
|
|
14
|
+
classifiers = [
|
|
15
|
+
"Intended Audience :: Information Technology",
|
|
16
|
+
"Intended Audience :: System Administrators",
|
|
17
|
+
"Operating System :: OS Independent",
|
|
18
|
+
"Programming Language :: Python :: 3 :: Only",
|
|
19
|
+
"Programming Language :: Python",
|
|
20
|
+
"Topic :: Software Development :: Libraries :: Python Modules",
|
|
21
|
+
"Topic :: Software Development :: Libraries",
|
|
22
|
+
"Topic :: Software Development",
|
|
23
|
+
"Typing :: Typed",
|
|
24
|
+
"Intended Audience :: Developers",
|
|
25
|
+
"License :: OSI Approved :: MIT License",
|
|
26
|
+
"Programming Language :: Python :: 3 :: Only",
|
|
27
|
+
"Programming Language :: Python :: 3.11",
|
|
28
|
+
"Programming Language :: Python :: 3.12",
|
|
29
|
+
"Programming Language :: Python :: 3.13",
|
|
30
|
+
]
|
|
31
|
+
dependencies = [
|
|
32
|
+
"boto3",
|
|
33
|
+
"click",
|
|
34
|
+
"click_spinner",
|
|
35
|
+
"coloredlogs",
|
|
36
|
+
"environs",
|
|
37
|
+
"loguru",
|
|
38
|
+
"paramiko",
|
|
39
|
+
"questionary",
|
|
40
|
+
"typing_extensions",
|
|
41
|
+
]
|
|
42
|
+
dynamic = ["version"]
|
|
43
|
+
|
|
44
|
+
[project.urls]
|
|
45
|
+
Homepage = "https://github.com/mvanderlee/aws-ssh-utils"
|
|
46
|
+
|
|
47
|
+
[project.scripts]
|
|
48
|
+
aws_ssh = "aws_ssh_utils.ssh:cli"
|
|
49
|
+
|
|
50
|
+
[project.optional-dependencies]
|
|
51
|
+
|
|
52
|
+
dev = [
|
|
53
|
+
"boto3-stubs[ec2,emr]",
|
|
54
|
+
]
|
|
55
|
+
test = [
|
|
56
|
+
]
|
|
57
|
+
publish = [
|
|
58
|
+
"hatch >= 1.7.0",
|
|
59
|
+
]
|
|
60
|
+
|
|
61
|
+
[tool.hatch.build]
|
|
62
|
+
include = [
|
|
63
|
+
"aws_ssh_utils",
|
|
64
|
+
]
|
|
65
|
+
|
|
66
|
+
[tool.hatch.version]
|
|
67
|
+
path = "aws_ssh_utils/__init__.py"
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
[tool.coverage.run]
|
|
71
|
+
parallel = true
|
|
72
|
+
source = [
|
|
73
|
+
"aws_ssh_utils"
|
|
74
|
+
]
|
|
75
|
+
|
|
76
|
+
[tool.autopep8]
|
|
77
|
+
max_line_length = 130
|
|
78
|
+
|
|
79
|
+
[tool.ruff]
|
|
80
|
+
exclude = [
|
|
81
|
+
"*.ipynb",
|
|
82
|
+
"**/_version.py",
|
|
83
|
+
]
|
|
84
|
+
lint.select = [
|
|
85
|
+
"E", # pycodestyle errors
|
|
86
|
+
"W", # pycodestyle warnings
|
|
87
|
+
"F", # pyflakes
|
|
88
|
+
"I", # isort
|
|
89
|
+
"C", # flake8-comprehensions
|
|
90
|
+
"B", # flake8-bugbear
|
|
91
|
+
"A", # flake8-builtins
|
|
92
|
+
"N", # pep8-naming
|
|
93
|
+
"UP", # pyupgrade
|
|
94
|
+
"S", # flake8-bandit
|
|
95
|
+
"COM", # flake8-commas
|
|
96
|
+
"ISC", # flake8-implicit-str-concat
|
|
97
|
+
"INP", # flake8-no-pep420
|
|
98
|
+
"PIE", # flake8-pie
|
|
99
|
+
"SIM", # flake8-simplify
|
|
100
|
+
"RUF", # Ruff-specific rules
|
|
101
|
+
"T20", # flake8-print
|
|
102
|
+
"PT", # flake8-pytest-style
|
|
103
|
+
]
|
|
104
|
+
lint.ignore = [
|
|
105
|
+
"B008", # do not perform function calls in argument defaults
|
|
106
|
+
"B028", # No explicit stacklevel argument found.
|
|
107
|
+
"C408", # unnecessary-collection-call - Rewrite as a literal
|
|
108
|
+
"C901", # too complex
|
|
109
|
+
"PT003", # pytest-extraneous-scope-function
|
|
110
|
+
|
|
111
|
+
# TODO: remove
|
|
112
|
+
"A002", # builtin-argument-shadowing
|
|
113
|
+
"PT011", # pytest-raises-too-broad
|
|
114
|
+
"N818", # error-suffix-on-exception-name
|
|
115
|
+
|
|
116
|
+
# Unsure, need more data
|
|
117
|
+
"SIM117", # multiple-with-statements - In general good, but sometimes becomes harder to read, i.e.: aiohttp?
|
|
118
|
+
]
|
|
119
|
+
# The formatter wraps lines at a length of 88.
|
|
120
|
+
line-length = 130 # Enforce 130. But ~100 is recommended.
|
|
121
|
+
|
|
122
|
+
[tool.ruff.lint.pycodestyle]
|
|
123
|
+
max-line-length = 150 # E501 reports lines that exceed the length of 150.
|
|
124
|
+
|
|
125
|
+
[tool.ruff.format]
|
|
126
|
+
quote-style = "preserve" # Don't change quotes - https://docs.astral.sh/ruff/settings/#format_quote-style
|
|
127
|
+
|
|
128
|
+
[tool.ruff.lint.per-file-ignores]
|
|
129
|
+
"__init__.py" = ["F401"]
|
|
130
|
+
|
|
131
|
+
[tool.ruff.lint.flake8-bugbear]
|
|
132
|
+
# Allow default arguments like, e.g., `data: List[str] = fastapi.Query(None)`.
|
|
133
|
+
extend-immutable-calls = []
|
|
134
|
+
|
|
135
|
+
[tool.pytest.ini_options]
|
|
136
|
+
# https://pytest-asyncio.readthedocs.io/en/latest/reference/configuration.html
|
|
137
|
+
asyncio_mode = "auto"
|
|
138
|
+
asyncio_default_test_loop_scope = "session"
|
|
139
|
+
asyncio_default_fixture_loop_scope = "session"
|
|
140
|
+
|
|
141
|
+
[tool.pyright]
|
|
142
|
+
venvPath = "."
|
|
143
|
+
venv = ".venv"
|
|
144
|
+
|
|
145
|
+
include = [
|
|
146
|
+
"aws_ssh_utils/",
|
|
147
|
+
]
|
|
148
|
+
|
|
149
|
+
pythonVersion = "3.11"
|
|
150
|
+
|
|
151
|
+
# basedpyright exclusions
|
|
152
|
+
reportAny = "none"
|
|
153
|
+
reportUnreachable = "none"
|
|
154
|
+
reportExplicitAny = "none"
|
|
155
|
+
reportUnusedCallResult = "none"
|
|
156
|
+
reportUnknownMemberType = "none"
|
|
157
|
+
reportUnusedParameter = "none"
|
|
158
|
+
reportUntypedFunctionDecorator = "none"
|
|
159
|
+
reportMissingTypeStubs = "none"
|
|
160
|
+
reportUnknownVariableType = "none"
|
|
161
|
+
reportImplicitStringConcatenation = "none"
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "0.1.1"
|
|
@@ -1,102 +0,0 @@
|
|
|
1
|
-
[build-system]
|
|
2
|
-
requires = ["hatchling"]
|
|
3
|
-
build-backend = "hatchling.build"
|
|
4
|
-
|
|
5
|
-
[project]
|
|
6
|
-
name = "aws-ssh-utils"
|
|
7
|
-
description = "Easy AWS SSHing"
|
|
8
|
-
readme = "README.md"
|
|
9
|
-
requires-python = ">=3.11"
|
|
10
|
-
license = "MIT"
|
|
11
|
-
authors = [
|
|
12
|
-
{ name = "Michiel Vanderlee", email = "jmt.vanderlee@gmail.com" },
|
|
13
|
-
]
|
|
14
|
-
classifiers = [
|
|
15
|
-
"Intended Audience :: Information Technology",
|
|
16
|
-
"Intended Audience :: System Administrators",
|
|
17
|
-
"Operating System :: OS Independent",
|
|
18
|
-
"Programming Language :: Python :: 3 :: Only",
|
|
19
|
-
"Programming Language :: Python",
|
|
20
|
-
"Topic :: Software Development :: Libraries :: Python Modules",
|
|
21
|
-
"Topic :: Software Development :: Libraries",
|
|
22
|
-
"Topic :: Software Development",
|
|
23
|
-
"Typing :: Typed",
|
|
24
|
-
"Intended Audience :: Developers",
|
|
25
|
-
"License :: OSI Approved :: MIT License",
|
|
26
|
-
"Programming Language :: Python :: 3 :: Only",
|
|
27
|
-
"Programming Language :: Python :: 3.11",
|
|
28
|
-
"Programming Language :: Python :: 3.12",
|
|
29
|
-
"Programming Language :: Python :: 3.13",
|
|
30
|
-
]
|
|
31
|
-
dependencies = [
|
|
32
|
-
"boto3",
|
|
33
|
-
"click",
|
|
34
|
-
"click_spinner",
|
|
35
|
-
"coloredlogs",
|
|
36
|
-
"environs",
|
|
37
|
-
"loguru",
|
|
38
|
-
"paramiko",
|
|
39
|
-
"questionary",
|
|
40
|
-
]
|
|
41
|
-
dynamic = ["version"]
|
|
42
|
-
|
|
43
|
-
[project.urls]
|
|
44
|
-
Homepage = "https://github.com/mvanderlee/aws_ssh_utils"
|
|
45
|
-
|
|
46
|
-
[project.scripts]
|
|
47
|
-
aws_ssh = "aws_ssh_utils.ssh:cli"
|
|
48
|
-
|
|
49
|
-
[project.optional-dependencies]
|
|
50
|
-
|
|
51
|
-
dev = [
|
|
52
|
-
"boto3-stubs[ec2,emr]",
|
|
53
|
-
]
|
|
54
|
-
test = [
|
|
55
|
-
]
|
|
56
|
-
publish = [
|
|
57
|
-
"hatch >= 1.7.0",
|
|
58
|
-
]
|
|
59
|
-
|
|
60
|
-
[tool.hatch.build]
|
|
61
|
-
include = [
|
|
62
|
-
"aws_ssh_utils",
|
|
63
|
-
]
|
|
64
|
-
|
|
65
|
-
[tool.hatch.version]
|
|
66
|
-
path = "aws_ssh_utils/__init__.py"
|
|
67
|
-
|
|
68
|
-
[tool.isort]
|
|
69
|
-
profile = "hug"
|
|
70
|
-
line_length = 100
|
|
71
|
-
|
|
72
|
-
[tool.coverage.run]
|
|
73
|
-
parallel = true
|
|
74
|
-
source = [
|
|
75
|
-
"aws_ssh_utils"
|
|
76
|
-
]
|
|
77
|
-
|
|
78
|
-
[tool.ruff]
|
|
79
|
-
select = [
|
|
80
|
-
"E", # pycodestyle errors
|
|
81
|
-
"W", # pycodestyle warnings
|
|
82
|
-
"F", # pyflakes
|
|
83
|
-
# "I", # isort
|
|
84
|
-
"C", # flake8-comprehensions
|
|
85
|
-
"B", # flake8-bugbear
|
|
86
|
-
]
|
|
87
|
-
ignore = [
|
|
88
|
-
"B008", # do not perform function calls in argument defaults
|
|
89
|
-
"B028", # No explicit stacklevel argument found.
|
|
90
|
-
"C901", # too complex
|
|
91
|
-
]
|
|
92
|
-
line-length = 130 # Enforce 130. But ~100 is recommended.
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
[tool.ruff.per-file-ignores]
|
|
96
|
-
"__init__.py" = ["F401"]
|
|
97
|
-
"migrations/env.py" = ["E402"]
|
|
98
|
-
|
|
99
|
-
[tool.pytest.ini_options]
|
|
100
|
-
# https://pytest-asyncio.readthedocs.io/en/latest/reference/configuration.html
|
|
101
|
-
asyncio_mode = "auto"
|
|
102
|
-
asyncio_default_fixture_loop_scope = "session"
|
|
File without changes
|