aws-ssh-utils 0.2.3__tar.gz → 0.2.4__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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: aws-ssh-utils
3
- Version: 0.2.3
3
+ Version: 0.2.4
4
4
  Summary: Easy AWS SSHing
5
5
  Project-URL: Homepage, https://github.com/mvanderlee/aws-ssh-utils
6
6
  Author-email: Michiel Vanderlee <jmt.vanderlee@gmail.com>
@@ -0,0 +1 @@
1
+ __version__ = "0.2.4"
@@ -6,6 +6,7 @@ import sys
6
6
  import textwrap
7
7
  from dataclasses import dataclass, field
8
8
  from hashlib import sha256
9
+ from pathlib import Path
9
10
  from types import TracebackType
10
11
  from typing import TYPE_CHECKING, Any
11
12
 
@@ -85,10 +86,14 @@ def cli(
85
86
  **kwargs: Any,
86
87
  ):
87
88
  log_format = (
88
- "<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</green> | "
89
- "<level>{level: <8}</level> | "
90
- "<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>"
91
- ) if long_log else "<level>{message}</level>"
89
+ (
90
+ "<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</green> | "
91
+ "<level>{level: <8}</level> | "
92
+ "<cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>"
93
+ )
94
+ if long_log
95
+ else "<level>{message}</level>"
96
+ )
92
97
  level = "DEBUG" if verbose else 100 if quiet else "INFO"
93
98
 
94
99
  logger.remove()
@@ -110,20 +115,25 @@ def ec2_ssh(
110
115
  **kwargs: Any,
111
116
  ):
112
117
  '''
113
- Asks user which EC2 instance they want to connect to,
114
- then opens an interactive SSH session to the instance
118
+ Asks user which EC2 instance they want to connect to,
119
+ then opens an interactive SSH session to the instance
115
120
 
116
121
  '''
117
122
  try:
118
123
  b3s = boto3.Session(profile_name=profile, region_name=region)
119
124
 
120
125
  ip, user, key_file, instance_name = get_ec2_ssh_options(
121
- b3s=b3s, user=user, use_private_ip=private, key_file=key_file,
126
+ b3s=b3s,
127
+ user=user,
128
+ use_private_ip=private,
129
+ key_file=key_file,
122
130
  )
123
131
  terminal_title = f'{user}@{instance_name}'
124
132
 
125
133
  SSHShell(
126
- hostname=ip, username=user, key_filename=key_file,
134
+ hostname=ip,
135
+ username=user,
136
+ key_filename=key_file,
127
137
  terminal_title=terminal_title,
128
138
  ).connect()
129
139
  except ShellError as e:
@@ -132,7 +142,8 @@ def ec2_ssh(
132
142
  except ClientError as e:
133
143
  if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
134
144
  logger.log(
135
- logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
145
+ logging.CRITICAL,
146
+ 'Your AWS Token has expired. Please update and try again.',
136
147
  )
137
148
  exit(1)
138
149
 
@@ -152,8 +163,8 @@ def emr_ssh(
152
163
  **kwargs: Any,
153
164
  ):
154
165
  '''
155
- Asks user which Cluster and EC2 instance they want to connect to,
156
- then opens an interactive SSH session to the instance
166
+ Asks user which Cluster and EC2 instance they want to connect to,
167
+ then opens an interactive SSH session to the instance
157
168
  '''
158
169
  try:
159
170
  b3s = boto3.Session(profile_name=profile, region_name=region)
@@ -166,7 +177,8 @@ def emr_ssh(
166
177
 
167
178
  if key_file is None:
168
179
  key_file = get_emr_ssh_key_file_for_cluster(
169
- emr, emr_instance.cluster_id,
180
+ emr,
181
+ emr_instance.cluster_id,
170
182
  )
171
183
 
172
184
  if private:
@@ -199,7 +211,8 @@ def emr_ssh(
199
211
  except ClientError as e:
200
212
  if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
201
213
  logger.log(
202
- logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
214
+ logging.CRITICAL,
215
+ 'Your AWS Token has expired. Please update and try again.',
203
216
  )
204
217
  exit(1)
205
218
 
@@ -219,18 +232,15 @@ def emr_ssh_all(
219
232
  **kwargs: Any,
220
233
  ):
221
234
  '''
222
- Asks user which Cluster and EC2 instance they want to connect to,
223
- Then prints a tmux cli statement that will open a new session
224
- with a window per ec2 instance with ssh shell already opened.
235
+ Asks user which Cluster and EC2 instance they want to connect to,
236
+ Then prints a tmux cli statement that will open a new session
237
+ with a window per ec2 instance with ssh shell already opened.
225
238
  '''
226
239
  try:
227
240
  b3s = boto3.Session(profile_name=profile, region_name=region)
228
241
  emr = b3s.client('emr')
229
242
  cluster_id, cluster_name = prompt_for_emr_cluster(emr)
230
- grouped_instances = {
231
- k.split(" ")[0]: v
232
- for k, v in get_emr_instance_ips(emr, cluster_id).items()
233
- }
243
+ grouped_instances = {k.split(" ")[0]: v for k, v in get_emr_instance_ips(emr, cluster_id).items()}
234
244
 
235
245
  if user is None:
236
246
  user = 'hadoop'
@@ -281,7 +291,8 @@ def emr_ssh_all(
281
291
  except ClientError as e:
282
292
  if e.response.get('Error') and e.response.get('Error', {}).get('Code') == 'ExpiredTokenException':
283
293
  logger.log(
284
- logging.CRITICAL, 'Your AWS Token has expired. Please update and try again.',
294
+ logging.CRITICAL,
295
+ 'Your AWS Token has expired. Please update and try again.',
285
296
  )
286
297
  exit(1)
287
298
 
@@ -299,13 +310,18 @@ class SSHShell:
299
310
 
300
311
  def __post_init__(self):
301
312
  if self.key_filename is not None:
302
- if not os.path.isfile(self.key_filename):
313
+ key_filename = Path(self.key_filename)
314
+ if not key_filename.exists():
303
315
  raise ValueError(f'File {self.key_filename} does not exist')
304
316
 
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)
317
+ self.private_key = paramiko.PKey.from_path(key_filename)
318
+ # Load first public key - opkssh used to postfix with ".pub", but now uses "-cert.pub"
319
+ for postfix in (".pub", "-cert.pub"):
320
+ stem_postfix, suffix = postfix.split('.', maxsplit=1)
321
+ public_key_path = key_filename.with_stem(f"{key_filename.stem}{stem_postfix}").with_suffix(f".{suffix}")
322
+ if public_key_path.exists() and public_key_path.is_file():
323
+ self.private_key.load_certificate(str(public_key_path.absolute()))
324
+ break
309
325
 
310
326
  def connect(self):
311
327
  self._open()
@@ -313,9 +329,7 @@ class SSHShell:
313
329
  self._close()
314
330
 
315
331
  def _open(self) -> paramiko.Channel:
316
- logger.info(
317
- f'Opening SSH: ssh -i {self.key_filename} {self.username}@{self.hostname}',
318
- )
332
+ logger.info(f'Opening SSH: ssh -i {self.key_filename} {self.username}@{self.hostname}')
319
333
 
320
334
  terminal_size = os.get_terminal_size()
321
335
 
@@ -355,7 +369,8 @@ class SSHShell:
355
369
  channel = ssh_client.get_transport().open_session() # pyright: ignore[reportOptionalMemberAccess]
356
370
  channel.get_pty(
357
371
  term=os.getenv('TERM', 'xterm-256color'),
358
- width=terminal_size.columns, height=terminal_size.lines,
372
+ width=terminal_size.columns,
373
+ height=terminal_size.lines,
359
374
  )
360
375
  channel.invoke_shell()
361
376
 
@@ -368,7 +383,8 @@ class SSHShell:
368
383
 
369
384
  def _launch(self):
370
385
  interactive_shell(
371
- self._channel, allow_title_changes=not self.terminal_title,
386
+ self._channel,
387
+ allow_title_changes=not self.terminal_title,
372
388
  )
373
389
 
374
390
  def _close(self):
@@ -380,7 +396,12 @@ class SSHShell:
380
396
  def __enter__(self):
381
397
  return self._open()
382
398
 
383
- def __exit__(self, exception_type: type[BaseException] | None, exception_value: BaseException | None, traceback: TracebackType | None):
399
+ def __exit__(
400
+ self,
401
+ exception_type: type[BaseException] | None,
402
+ exception_value: BaseException | None,
403
+ traceback: TracebackType | None,
404
+ ):
384
405
  self._launch()
385
406
  self._close()
386
407
 
@@ -415,10 +436,7 @@ def get_ec2_ssh_options(
415
436
  key_file: str | None = None,
416
437
  ) -> tuple[str, str, str | None, str]:
417
438
  instance = prompt_for_ec2_instance(b3s)
418
- instance_tags = {
419
- tag.get('Key', '').lower(): tag.get('Value', '')
420
- for tag in instance.tags
421
- }
439
+ instance_tags = {tag.get('Key', '').lower(): tag.get('Value', '') for tag in instance.tags}
422
440
 
423
441
  if key_file is None:
424
442
  # Support https://github.com/openpubkey/opkssh via tags
@@ -426,10 +444,12 @@ def get_ec2_ssh_options(
426
444
  key_file = get_opkssh_key_file(instance_tags[OPKSSH_PROVIDER_TAG])
427
445
  elif OPKSSH_ISSUER_TAG in instance_tags and OPKSSH_CLIENT_TAG in instance_tags:
428
446
  key_file = get_opkssh_key_file(
429
- ','.join([
430
- instance_tags[OPKSSH_ISSUER_TAG],
431
- instance_tags[OPKSSH_CLIENT_TAG],
432
- ]),
447
+ ','.join(
448
+ [
449
+ instance_tags[OPKSSH_ISSUER_TAG],
450
+ instance_tags[OPKSSH_CLIENT_TAG],
451
+ ],
452
+ ),
433
453
  )
434
454
  else:
435
455
  key_file = try_to_find_ssh_key_file(instance.key_name)
@@ -457,10 +477,7 @@ def get_emr_ssh_options(
457
477
 
458
478
  instance_ips = sorted(instance_ips, key=lambda ips: ips.private or "")
459
479
 
460
- instance_options = [
461
- f'{ip.private} ({ip.public})' if ip.public else f'{ip.private}'
462
- for ip in instance_ips
463
- ]
480
+ instance_options = [f'{ip.private} ({ip.public})' if ip.public else f'{ip.private}' for ip in instance_ips]
464
481
  instance_ip = questionary.select(
465
482
  'Which instance do you want to connect to?',
466
483
  choices=instance_options,
@@ -476,10 +493,7 @@ def get_emr_ssh_key_file_for_cluster(
476
493
  ) -> str | None:
477
494
  with click_spinner.spinner():
478
495
  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
- }
496
+ cluster_tags = {tag.get('Key', '').lower(): tag.get('Value', '') for tag in cluster['Cluster'].get('Tags', [])}
483
497
 
484
498
  key_file = key_name = None
485
499
  # Support https://github.com/openpubkey/opkssh via tags
@@ -487,16 +501,16 @@ def get_emr_ssh_key_file_for_cluster(
487
501
  key_file = get_opkssh_key_file(cluster_tags[OPKSSH_PROVIDER_TAG])
488
502
  elif OPKSSH_ISSUER_TAG in cluster_tags and OPKSSH_CLIENT_TAG in cluster_tags:
489
503
  key_file = get_opkssh_key_file(
490
- ','.join([
491
- cluster_tags[OPKSSH_ISSUER_TAG],
492
- cluster_tags[OPKSSH_CLIENT_TAG],
493
- ]),
504
+ ','.join(
505
+ [
506
+ cluster_tags[OPKSSH_ISSUER_TAG],
507
+ cluster_tags[OPKSSH_CLIENT_TAG],
508
+ ],
509
+ ),
494
510
  )
495
511
 
496
512
  if key_file is None:
497
- key_name = cluster["Cluster"].get(
498
- "Ec2InstanceAttributes", {},
499
- ).get("Ec2KeyName", "")
513
+ key_name = cluster["Cluster"].get("Ec2InstanceAttributes", {}).get("Ec2KeyName", "")
500
514
  key_file = try_to_find_ssh_key_file(key_name)
501
515
 
502
516
  if key_file is None:
@@ -525,15 +539,17 @@ def prompt_for_ec2_instance(
525
539
  prompt: str = 'Which EC2 instance do you want to connect to?',
526
540
  ) -> "Instance":
527
541
  '''
528
- Discovers the runnin EC2 clusters and asks the user which to use.
542
+ Discovers the runnin EC2 clusters and asks the user which to use.
529
543
 
530
- Returns the EC2 object that the user selected.
544
+ Returns the EC2 object that the user selected.
531
545
  '''
532
546
  running_instances = get_running_ec2_instances(b3s)
533
547
 
534
- name_contains = questionary.text(
535
- 'Provide a name filter or leave blank to show all',
536
- ).unsafe_ask().lower()
548
+ name_contains = (
549
+ questionary.text('Provide a name filter or leave blank to show all') #
550
+ .unsafe_ask()
551
+ .lower()
552
+ )
537
553
 
538
554
  with click_spinner.spinner():
539
555
  grouped_by_name = {
@@ -565,15 +581,23 @@ def get_ec2_name(ec2_instance: "Instance") -> str:
565
581
 
566
582
 
567
583
  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
- )
584
+ opkssh_key_file = Path.home() / '.ssh/' / f'opkssh_{sha256(provider.encode()).hexdigest()}'
573
585
 
574
- if os.path.exists(opkssh_key_file) and os.path.getmtime(opkssh_key_file) > dt.datetime.now().timestamp() - 86400:
586
+ if opkssh_key_file.exists() and opkssh_key_file.stat().st_mtime > dt.datetime.now().timestamp() - 86400:
575
587
  logger.info('Found existing opkssh key')
576
588
  else:
589
+ if opkssh_key_file.exists() and opkssh_key_file.is_file():
590
+ logger.info("Detected expired opkssh key, deleting")
591
+ # Delete public keys - opkssh used to postfix with ".pub", but now uses "-cert.pub"
592
+ for postfix in (".pub", "-cert.pub"):
593
+ stem_postfix, suffix = postfix.split('.', maxsplit=1)
594
+ public_key_path = opkssh_key_file.with_stem(f"{opkssh_key_file.stem}{stem_postfix}").with_suffix(f".{suffix}")
595
+ if public_key_path.exists() and public_key_path.is_file():
596
+ public_key_path.unlink()
597
+
598
+ # Delete private key
599
+ opkssh_key_file.unlink()
600
+
577
601
  logger.info('Detected opkssh, logging in')
578
602
  cmd = [
579
603
  "opkssh",
@@ -581,7 +605,7 @@ def get_opkssh_key_file(provider: str) -> str:
581
605
  "--provider",
582
606
  provider,
583
607
  "-i",
584
- opkssh_key_file,
608
+ str(opkssh_key_file.absolute()),
585
609
  ]
586
610
  process = subprocess.Popen(cmd, shell=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE) # noqa: S603
587
611
  process.wait()
@@ -593,7 +617,7 @@ def get_opkssh_key_file(provider: str) -> str:
593
617
  logger.error(process.stderr.read().decode())
594
618
  sys.exit(process.returncode)
595
619
 
596
- return opkssh_key_file
620
+ return str(opkssh_key_file.absolute())
597
621
 
598
622
 
599
623
  class ConfirmAddPolicy(paramiko.client.MissingHostKeyPolicy):
@@ -612,9 +636,7 @@ class ConfirmAddPolicy(paramiko.client.MissingHostKeyPolicy):
612
636
  ).unsafe_ask()
613
637
  if should_add:
614
638
  client.get_host_keys().add(hostname, key.get_name(), key)
615
- host_key_filename: str | None = (
616
- client._host_keys_filename # pyright: ignore[reportAttributeAccessIssue]
617
- )
639
+ host_key_filename: str | None = client._host_keys_filename # pyright: ignore[reportAttributeAccessIssue]
618
640
  if host_key_filename is not None:
619
641
  client.save_host_keys(host_key_filename)
620
642
  logger.info('Added host key')
@@ -1 +0,0 @@
1
- __version__ = "0.2.3"
File without changes
File without changes
File without changes